feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+58
View File
@@ -0,0 +1,58 @@
# Race, Go. Dispatched by hand, and never a gate on a push or a tag: the release tag is
# cut only after `just gates` has already raced the tree, so this workflow is the
# explicit second opinion, not a step of the release.
#
# The race detector roughly doubles both time and memory, which the shared runner box
# cannot afford on every push. Locally it belongs to `just gates`, which runs it once
# per task; here it is an explicit decision rather than a routine.
#
# The matrix keeps the libm check the push pipeline once carried: the same linux/amd64
# oracle digests run against glibc (fedora, openeuler) and musl (alpine), which is
# exactly where floating-point kernels can drift. Dispatched, because three full
# sweeps are not affordable on every push.
#
# Every step is one command, so the step that fails is the gate that failed.
name: Race
on:
workflow_dispatch:
env:
# One core: parallelism buys no speed here and costs memory the box does not have.
GOFLAGS: -p=1
GOMAXPROCS: "2"
jobs:
race:
runs-on: ${{ matrix.runner }}
timeout-minutes: 45
strategy:
fail-fast: false
matrix:
include:
- runner: fedora
packages: dnf install -y git gcc perl
- runner: alpine
packages: apk add --no-cache git gcc perl musl-dev
- runner: openeuler
packages: dnf install -y git gcc perl
steps:
- name: Install git, gcc and Perl
# The race detector needs cgo, hence gcc; alpine adds musl-dev for the same
# reason. The installs are no-ops where the packages already exist.
run: ${{ matrix.packages }}
- uses: actions/checkout@v7
- uses: actions/setup-go@v6
with:
go-version-file: go.mod
cache: true
- name: Oracle digests for this platform
run: go test -run TestOracle -v .
- name: Race
# Equal to `packages` in the project's justfile: the logic packages,
# the main-program examples aside.
run: go test -race -count=1 -timeout 30m . ./internal/... ./grad/... ./integrate/... ./io/... ./linalg/... ./optim/... ./signal/... ./spmd/... ./stats/...
+183
View File
@@ -0,0 +1,183 @@
# Release, Go library. Runs on version tags (v1.2.3) pushed to main.
#
# A library ships no build assets, so there is no build matrix and no smoke test: the
# gate set minus race runs once at the tag, then the release is created from the
# matching CHANGELOG section. Race never runs on a push path or a tag; the local gate
# raced this tree before the tag was cut. Nothing is injected; the toolchain records
# the tag into the module's build information because the build simply happens there.
# The version contract these steps implement is in the `release` skill.
#
# Every scripted step is Perl with builtins only, and Perl drives curl through a list,
# so no argument is ever word-split, globbed or quoted wrong.
name: Release
on:
push:
tags: ["v*"]
env:
# One core: parallelism buys no speed here and costs memory the box does not have.
GOFLAGS: -p=1
GOMAXPROCS: "2"
jobs:
gates:
runs-on: fedora
timeout-minutes: 10
steps:
- name: Install git and Perl
# Both are no-ops where present. gcc existed for the race detector,
# which no longer runs in this pipeline.
run: dnf install -y git perl
- uses: actions/checkout@v7
- uses: actions/setup-go@v6
with:
go-version-file: go.mod
cache: true
- name: Validate the tag
env:
VERSION: ${{ gitea.ref_name }}
run: |
perl -e '
my $v = $ENV{VERSION} // q{};
$v =~ m{^v[0-9]+(\.[0-9]+){0,2}([-+].*)?$}
or die qq{ERROR: expected a semver tag like v1.2.3, got: $v\n};
print qq{tag $v\n};
'
- name: Format
run: |
perl -e '
open(my $g, q{-|}, q{gofmt}, q{-l}, q{.}) or die qq{gofmt: $!};
my @bad = <$g>;
close($g);
print @bad;
exit(@bad ? 1 : 0);
'
- name: Vet
run: go vet ./...
- name: Modernise
run: go fix -diff ./...
- name: Build
run: go build ./...
- name: Tests
# Equal to `packages` in the project's justfile, so the floor is the same
# number the local gate reports.
run: go test -count=1 -timeout 10m -coverprofile=coverage.out . ./internal/... ./grad/... ./integrate/... ./io/... ./linalg/... ./optim/... ./signal/... ./stats/...
- name: Coverage floor
run: |
perl -e '
open(my $c, q{-|}, q{go}, q{tool}, q{cover}, q{-func=coverage.out}) or die qq{cover: $!};
my $total;
while (my $l = <$c>) { $total = $1 if $l =~ m{^total:\s+\S+\s+([0-9.]+)%} }
close($c);
die qq{no total line in coverage.out\n} unless defined $total;
printf qq{Total coverage: %s%%\n}, $total;
exit($total < 80 ? 1 : 0);
'
release:
runs-on: fedora
timeout-minutes: 15
needs: gates
permissions:
# contents: read is required for the checkout: a job that declares any
# permissions gets a token scoped to exactly those, and releases: write
# alone leaves the fetch with no read access, which Gitea answers with
# a 404 "Repository not found". Verified on the instance 2026-09-16.
contents: read
releases: write
steps:
- name: Install Perl
run: dnf install -y perl
- uses: actions/checkout@v7
- name: Extract the CHANGELOG section
env:
VERSION: ${{ gitea.ref_name }}
run: |
# Each step derives what it needs from the tag, so no value has to travel
# between jobs.
perl -e '
my $v = $ENV{VERSION} // q{};
$v =~ s{^v}{};
open(my $vout, q{>}, q{version-no-v.txt}) or die qq{version-no-v.txt: $!};
print $vout $v;
close($vout);
open(my $in, q{<}, q{CHANGELOG.md}) or die qq{CHANGELOG.md: $!};
my @lines = <$in>;
close($in);
my ($start, $end) = (-1, scalar @lines);
for my $i (0 .. $#lines) {
if ($start < 0) { $start = $i if $lines[$i] =~ m{^##\s+\[\Q$v\E\]} }
elsif ($lines[$i] =~ m{^##\s+\[}) { $end = $i; last }
}
$start >= 0 or die qq{ERROR: no CHANGELOG section for $v, expected a heading like: ## [$v] - YYYY-MM-DD\n};
my @body = grep { m{\S} } @lines[$start + 1 .. $end - 1];
@body or die qq{ERROR: the CHANGELOG section for $v is empty\n};
open(my $out, q{>}, q{release-body.md}) or die qq{release-body.md: $!};
print $out @body;
close($out);
printf qq{notes for %s: %d lines\n}, $v, scalar @body;
'
- name: Build the release request
run: |
perl -e '
open(my $vin, q{<}, q{version-no-v.txt}) or die qq{version-no-v.txt: $!};
my $v = <$vin>;
close($vin);
chomp $v;
open(my $in, q{<:raw}, q{release-body.md}) or die qq{release-body.md: $!};
my $body = do { local $/; <$in> };
close($in);
# Byte-oriented escaping: JSON is UTF-8, so non-ASCII passes through and
# only the characters JSON forbids are rewritten.
$body =~ s/([\\"])/\\$1/g;
$body =~ s/\t/\\t/g;
$body =~ s/\r//g;
$body =~ s/\n/\\n/g;
$body =~ s/([\x00-\x08\x0b\x0c\x0e-\x1f])/sprintf(q{\u%04x}, ord($1))/ge;
my $json = sprintf(qq{{"tag_name":"v%s","name":"v%s","body":"%s","draft":false,"prerelease":false}}, $v, $v, $body);
open(my $out, q{>}, q{release.json}) or die qq{release.json: $!};
print $out $json;
close($out);
print qq{release.json written for v$v\n};
'
- name: Create the release
env:
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
GITEA_SERVER_URL: ${{ gitea.server_url }}
GITEA_REPOSITORY: ${{ gitea.repository }}
VERSION: ${{ gitea.ref_name }}
run: |
perl -e '
my @cmd = (q{curl}, q{-sS}, q{-o}, q{response.json}, q{-w}, q{%{http_code}},
q{-H}, qq{Authorization: token $ENV{GITEA_TOKEN}},
q{-H}, q{Content-Type: application/json},
q{-X}, q{POST},
qq{$ENV{GITEA_SERVER_URL}/api/v1/repos/$ENV{GITEA_REPOSITORY}/releases},
q{--data-binary}, q{@release.json});
open(my $curl, q{-|}, @cmd) or die qq{curl: $!};
my $code = <$curl>;
my $ok = close($curl);
my $exit = $? >> 8;
$code = defined $code ? $code : q{};
$ok or die qq{ERROR: curl failed (exit $exit) calling $ENV{GITEA_SERVER_URL}\n};
open(my $r, q{<:raw}, q{response.json}) or die qq{response.json: $!};
my $body = do { local $/; <$r> };
close($r);
$code eq q{201} or die qq{ERROR: the release was not created, HTTP $code: $body\n};
$body =~ m{"id"\s*:\s*([0-9]+)} or die qq{ERROR: no release id in the response: $body\n};
print qq{release v$ENV{VERSION} is live (id $1)\n};
'
+109
View File
@@ -0,0 +1,109 @@
# Test, Go. Push and pull request to development. Never on main.
#
# The gates are the ones the justfile's `gates` recipe runs, minus race: the shared
# runner box cannot afford the race detector on every push, so it lives in race.yml,
# dispatched by hand. The box is small and sits beside Gitea, so parallelism is bounded
# on purpose and everything runs in one job; extra jobs would duplicate the checkout,
# the Go setup and the dependency download without buying any parallelism.
#
# Every step is one command, so the step that fails is the gate that failed, and no shell
# option has to be trusted for the run to stop. The scripted steps are Perl, not shell and
# not Python: Perl behaves the same on both runner images, there is no bashism to trip over
# on ash, and it is one language instead of two. The Perl uses builtins only, because
# nothing beyond `perl` itself may be assumed present.
#
# Project facts the template had to bend for: the portable build runs at the toolchain
# default with nothing pinned, the oracle digests are recorded per platform there, and
# the package pattern is the logic packages of the project's justfile `packages`,
# which keeps the main-program examples out of the suite. The benchmark smoke that
# once rode along here is retired outright: the minimum degree battery's 3-D mesh
# scan alone runs for minutes on one core and allocates terabytes cumulatively, so
# no form of it fits the shared box, and benchmarking is deliberate work on a
# developer machine, where the battery is survivable and the numbers are the point.
name: Test
on:
push:
branches: [development]
pull_request:
branches: [development]
env:
# One core: parallelism buys no speed here and costs memory the box does not have.
GOFLAGS: -p=1
GOMAXPROCS: "2"
# A superseded run of the same ref is cancelled instead of queueing behind one that
# no longer matters. Verified on Gitea 1.27.1 on 2026-09-17: a queued run whose ref
# moved on is cancelled before it ever reaches the runner, while a run already
# dispatched there runs to completion.
concurrency:
group: ${{ gitea.workflow }}-${{ gitea.ref }}
cancel-in-progress: true
jobs:
test:
runs-on: fedora
timeout-minutes: 10
steps:
- name: Install git and Perl
# The runner images are minimal: checkout needs git and the scripted steps
# below are Perl. Both installs are no-ops where the package already exists.
run: dnf install -y git perl
- uses: actions/checkout@v7
- uses: actions/setup-go@v6
with:
# The module is the source of truth for the version, so it cannot drift.
go-version-file: go.mod
cache: true
# The steps follow the `gates` order of the justfile contract: build, format,
# vet, test. The vet gate is go vet and go fix -diff, two steps here.
- name: Build
# The examples are main programs; the build is what compiles them.
run: go build ./...
- name: Format
run: |
perl -e '
open(my $g, q{-|}, q{gofmt}, q{-l}, q{.}) or die qq{gofmt: $!};
my @bad = <$g>;
close($g);
print @bad;
exit(@bad ? 1 : 0);
'
- name: Vet
run: go vet ./...
- name: Modernise
# Exits non-zero when it has something to rewrite, so it needs no output capture.
run: go fix -diff ./...
- name: Tests
# Equal to `packages` in the project's justfile, so the floor is the same
# number the local gate reports. The inner timeout matches the job's, so a
# hanging test reports its own goroutine dump rather than a silent job kill.
# The local `just test` allows 30 minutes for a warm 32-core box; this
# runner is one shared core, where the suite stays well inside the ten
# minutes its budget has always allowed.
run: go test -count=1 -timeout 10m -coverprofile=coverage.out . ./internal/... ./grad/... ./integrate/... ./io/... ./linalg/... ./optim/... ./signal/... ./spmd/... ./stats/...
- name: Coverage floor
run: |
perl -e '
open(my $c, q{-|}, q{go}, q{tool}, q{cover}, q{-func=coverage.out}) or die qq{cover: $!};
my $total;
while (my $l = <$c>) { $total = $1 if $l =~ m{^total:\s+\S+\s+([0-9.]+)%} }
close($c);
die qq{no total line in coverage.out\n} unless defined $total;
printf qq{Total coverage: %s%%\n}, $total;
exit($total < 80 ? 1 : 0);
'
- name: Oracle digests for this platform
# TestOracle verifies the pinned digest block for GOOS/GOARCH and skips
# loudly with instructions when the platform has none yet.
run: go test -run TestOracle -v .
+5
View File
@@ -0,0 +1,5 @@
.idea/
.zcode/
bin/
coverage.out
*.test
+348
View File
@@ -0,0 +1,348 @@
# Changelog
All notable changes to **Tensor** are documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [1.0.0] - 2026-09-03
The initial release of Tensor, a scientific computing library in pure
Go: immutable n-dimensional arrays over a wide element-type set,
dense and sparse linear algebra, differential equations, quadrature,
signal transforms, statistics up to mixed models and hidden Markov
models, optimisation from simplex to stochastic global search, a
reverse-mode differentiable core, deterministic SVG plotting and an
experimental SPMD package, with no third-party dependencies and a
deterministic, parallel execution model.
The module is one core package plus one package per domain:
`sourcedock.dev/petrbalvin/tensor` carries the `Array` core with
element-wise math, special functions and the reproducible generator,
re-exporting the whole surface of every domain package; `linalg` the
dense and sparse solvers; `signal` the transforms, filters and
stencils; `integrate` the differential equations and quadrature;
`stats` the distributions and inference; `optim` the fitting and root
finding; `io` the CSV, FITS, HDF5, NetCDF and memory-mapped readers
and writers; `grad` the differentiable core; `plot` the deterministic
figures. Import what you use: the domain packages depend on the core,
never on each other except through five one-directional edges, and
nothing below the root imports the root. The experimental `spmd`
package stands beside them, imported explicitly: a distributed
program names it, the facade does not.
### Added
**Core arrays.**
- Immutable, shape-checked n-dimensional arrays over int64, float32,
float64, complex128, IEEE 754 float16 and the narrow integer types
Int8, Uint8, Int16, Uint16, Int32 and Uint32, with Bool beside
them, under a strict promotion ladder, row-major layout and
multi-rank text formatting. Constructors cover literals
(`FromInts`, `FromFloat32s`, `FromFloats`, `FromComplexes`,
`FromFloat16s`, `FromInt8s` and the family around them), filled and
ranged builders (`Zeros`, `Ones`, `FullI`/`FullF`/`FullC`, `Range`,
`RangeBy`, `Linspace`, `Grid`), byte loading and dtype conversions
(`WithInt`/`WithFloat`/`WithComplex`, and `Astype`, which
range-checks every conversion into a narrow target with an error
naming the value and its index).
- Element-wise arithmetic with scalar variants, transcendentals,
comparisons that answer a bool mask of one byte per element
composing through `And`, `Or`, `Xor` and `Not` (`Where` and
`Select` keep the int mask), and reductions from `Sum`, `Mean`,
`Min`/`Max`, `Prod` and `Dot` through axis variants, `ArgMax`,
`TopK`, `CumSum` and `CumProd` to `SumKahan` compensated summation.
Every float fold cuts its line through a canonical partition fixed
by the length alone and combines the partials through a balanced
tree, so a reduction's answer is a function of the data alone,
never of the machine, and `CumSum` carries a Neumaier compensation
term so a long prefix sheds no small addend.
- Shape operations (`Reshape`, `Flatten`, `Squeeze`, `Transpose`,
`TransposeAxes`, `MoveAxis`, `Pad`, `Tile`, `Repeat`, `Flip`, `Roll`,
`Diag`, triangular extractions), indexing (contiguous `Slice`
selections as read-only payload views, `Gather`, `Scatter`, `Take`,
`Nonzero`, `Argwhere`, `SearchSorted`), and the toolkit pieces
`Einsum`, `Unique`, `OneHot`, `CrossProduct`, `Sort`, `ArgSort`,
and the numeric `Jacobian` of a vector function by central
differences.
- `MatMul2D`, the parallel cache-friendly kernel; sparse matrices as
`SparseCOO`; interpolation (`Interpolate`, the Fritsch-Carlson
`InterpolateMonotone`, `Interpolate2D`, `InterpolateGrid`, natural
cubic splines); and special functions: the gamma and beta families,
error functions, Bessel of integer order and, through
`BesselJRealOrder`, of any real order, Airy and Fresnel families,
exponential integrals, orthogonal polynomials, spherical harmonics,
elliptic integrals and Jacobi functions, and `Cosm1` for
`cos(x) − 1` at the arguments where the direct subtraction has no
correct significant bit.
- Quasirandom sequences for Monte Carlo integration: `HaltonPoints`
and the base-2 digital `SobolPoints` (Joe and Kuo initialisation,
40 dimensions).
- Parallel execution across every core with a fixed reduction order,
`SetNumCPU` to pin the worker count, and pooled scratch buffers
whose borrow path zeroes the window.
- The reproducible `Generator`: xoshiro256++ seeded through
splitmix64, stable across Go releases, with uniform, normal and
truncated-normal draws, shuffles and permutations; `Splitmix64` and
`Substream` are exported for callers who seed their own streams.
The distribution draws of the `stats` package take the same
generator, so a seeded program replays exactly.
**Linear algebra.**
- Dense factorisations and solves: LU (`Solve`, `Inv`, `Det`), QR
beside `RRQR`, the rank-revealing column-pivoted factorisation with
`RRQRRank` and `SolveRRQR`, whose rank-deficient path answers the
minimum-norm solution, Cholesky with rank-one update and downdate,
`LeastSquares`, the tridiagonal and cyclic-tridiagonal solvers,
`Pinverse`, `MatrixRank`, `Cond`.
- Eigenproblems in the symmetric, Hermitian-complex, general and
generalised forms; `SVD` and `SVDComplex` in real and complex
arithmetic; `SchurComplex`; the matrix functions `MatrixExp`,
`MatrixSqrt` and `MatrixLog`.
- Regularised and truncated solves for ill-posed systems:
`SolveTruncated` (rank truncation of the singular spectrum) and
`SolveTikhonov` (Tikhonov damping through the SVD).
- Sparse direct factorisations: `CSCFromCOO` and the `SparseCSC` view
with `ToCSR`/`ToCSC` conversions, `NewSparseCholesky`, the sparse
Cholesky with the elimination tree, the natural, reverse
Cuthill-McKee and minimum-degree orderings and the rank-one
`Update` and `Downdate`, and `NewSparseLU`, the Gilbert-Peierls
left-looking elimination with partial pivoting. One factorisation
solves any number of right-hand sides.
- Sparse iterative methods on the CSR view: `SpSolve` (conjugate
gradient), `SpSolveBiCGSTAB`, `SpLSQR` and `SpLSMR` for
overdetermined systems, the ILU(0) preconditioner, the Lanczos
eigensolver `SpEigen` with its general Arnoldi form, and
`SpExpApply`, the Krylov action of a matrix exponential. The
complex side mirrors it for Hermitian positive-definite and general
non-Hermitian operators, the shape Helmholtz and electromagnetics
problems need.
- Polynomial fitting, roots through the companion matrix, and the
fluent `Pipeline` chain over the element-wise surface.
**Signal and transforms.**
- Fourier transforms of any length: `FFT`/`IFFT`, `FFT2`, `FFT3`,
`FFTN`/`IFFTN`, the real-input `RFFT`/`IRFFT` pair and `FFTFreq`.
- Cosine and sine transforms (orthonormal types I to IV), the
type-1 non-uniform FFT by Gaussian gridding, and the short-time
Fourier transform with window choice.
- Spectral estimation: `WelchPSD`, `Spectrogram`, `LombScargle` for
unevenly sampled data, and the spectral Poisson solves in periodic,
Dirichlet and Neumann boundaries.
- Sample-rate conversion: `Decimate` behind a Kaiser anti-alias
filter, rational `Resample` and exact band-limited
`ResampleFourier`; the Hilbert `AnalyticSignal` and `Envelope`; and
`Chirp`, the linear frequency sweep synthesised from the
closed-form phase at each sample.
- Filter design: Butterworth, Chebyshev, inverse Chebyshev and
elliptic (Cauer) responses in low-pass, high-pass, band-pass and
band-stop forms with explicit ripple and attenuation budgets,
applied through `FilterApply` and, zero phase, through `Filtfilt`.
- Correlation and wavelets: FFT-based `Autocorrelate` and
`CrossCorrelate`, `PartialAutocorrelate`, the Haar `DWT`/`IDWT` and
the Daubechies families db2 to db8, and the analytic `CWT` (Morlet
and Mexican hat).
- Convolutions, pooling and windows: `Conv1D`/`Conv2D`/`Conv3D` with
groups and dilation, `ConvTranspose2D`, the max, average, adaptive
and global pooling families, the `MedianFilter` and `RankFilter`
families in one and two dimensions, the `SavitzkyGolay` smoother,
the `Gradient1D`/`Laplacian` stencils, and the public window
catalogue `WindowHann` through `WindowBox`, each in the symmetric
and the periodic convention.
- Time-series estimation: `KalmanFilter`, `ExtendedKalmanFilter` and
`UnscentedKalmanFilter` with the filtered states, covariance
history, innovations and the summed log likelihood;
`EstimateAR` through Yule-Walker over the Levinson recursion,
`EstimateARMA` through Hannan-Rissanen innovations, `SelectARMA`
over a lag grid by information criterion, and `ARMASpectrum` for
the theoretical one-sided spectrum.
**Differential equations and quadrature.**
- Initial value problems: `IntegrateODE` (adaptive Dormand-Prince
4(5)) with path and step recording, `IntegrateBDF2` and
`IntegrateBDFVar` (variable order 1 to 5, VODE-style step and order
adaptation) for stiff systems, `IntegrateROS4`, the L-stable
Rosenbrock-Wanner solver, `IntegrateBackwardEuler`, `IntegrateRK4`,
event detection with direction filters, `IntegrateDAE` for
semi-explicit index-1 differential-algebraic systems in mass-matrix
form, and the symplectic `IntegrateVerlet` beside `IntegrateYoshida4`
and `IntegrateMidpoint` for separable and general Hamiltonians.
- Boundary values: `IntegrateBoundary` by damped shooting with the
root finder of `optim`, and `SolveBoundaryCollocation` by
three-point Lobatto IIIA collocation on an adaptively refined mesh.
- Quadrature: adaptive Gauss-Legendre `IntegrateFunction`, fixed-node
`GaussLegendreNodes`, `IntegrateND`, globally adaptive cubature
over hyperrectangles, and `IntegrateFilon` for the oscillatory
integrals of a smooth amplitude against a cosine or sine carrier.
- Turnkey PDE evolution: the heat equation by Crank-Nicolson in one
dimension and Peaceman-Rachford ADI in two, the wave equation by
velocity Verlet in one dimension and an explicit central stencil in
two, advection by the monotone upwind and Koren-limited fluxes and
their advection-diffusion combination, CFL enforced everywhere.
- Finite elements: structured and arbitrary triangular meshes in two
dimensions and tetrahedral meshes in three, with
`SolvePoissonFEM2D` and `SolvePoissonFEM3D`, the piecewise-linear
Poisson assemblies through the sparse direct factorisation, with
Dirichlet lifting and natural Neumann boundaries.
**Statistics.**
- Distributions: CDFs, quantiles and matched random draws for the
normal, exponential, gamma, chi-square, Student t, F and their
noncentral forms, Poisson, binomial, negative binomial, Weibull,
lognormal, Pareto and Dirichlet laws, built on the incomplete gamma
and beta functions.
- Descriptives: mean-free moments, `Median`, `Quantile`, histograms
in one and two dimensions, the robust `MedianAbsoluteDeviation` and
`TrimmedMean`, and the rolling windows.
- Inference: `WelchTTest`, `KolmogorovSmirnovTest`, `MannWhitneyU`,
one-way `ANOVAOneWay`, `ChiSquareGoodnessOfFit`, `BootstrapCI`, the
rank correlations `SpearmanRho` and `KendallTau`, the
multiple-testing corrections `Bonferroni`, `Holm` and
`BenjaminiHochberg`, and the contingency table analyses
`FisherExactTest`, `ChiSquareIndependence`, `McNemarTest` and
`CramersV`.
- Models: `LinearRegression` with standard errors, t-tests, p-values,
R² and the model F-test, `WeightedLinearRegression`,
`LogisticRegression` and `PoissonRegression` on the exact
likelihood with Wald inference, `HuberRegression`,
`TheilSenRegression` and `QuantileRegression` for the robust and
distribution-free fits, lasso and elastic net over a documented
regularisation path, `PCA`, `KMeans` with k-means++ seeding,
`GaussianMixture` selected over a component grid by
`GaussianMixtureBIC`, Gaussian-process regression over
squared-exponential, Matern 3/2 and 5/2 and periodic kernels,
multivariate normal densities and draws, Gaussian `KernelDensity`,
`LinearMixedModel`, the Gaussian linear mixed model with grouped
random effects estimated by residual maximum likelihood, the
discrete hidden Markov model with its `Forward`, `Smooth` and
`Viterbi` recursions and its Baum-Welch fit, and
`HierarchicalClustering`, the agglomerative dendrogram with the
single, complete, average, centroid and Ward linkages and the
`Dendrogram` cuts into flat clusters.
**Optimisation.**
- Local: `Minimise` (Nelder-Mead simplex), `MinimiseLBFGS` with box
bounds and a projected-gradient convergence measure, and
`LevenbergMarquardt` with an optional analytic Jacobian. The
`LevenbergMarquardtFit` form reports χ², a named `FitStatus` and,
on request, the parameter covariance and per-residual weights
through `Sigma`, and the least squares and system solvers take
`ParallelJacobian`, an explicit opt-in that spreads the
finite-difference columns across workers with bit-identical
answers.
- Constrained: `MinimiseConstrained`, the augmented Lagrangian over
the box, so equality and inequality rows of `LinearConstraints`
compose with the walls, and `MinimiseNonlinearConstrained` for rows
that are arbitrary functions.
- Global: `MinimiseDifferentialEvolution` for multimodal,
derivative-free landscapes, `MinimiseCMAES` (the rank-one and
rank-mu update set, seeded through the generator) and
`MinimiseSimulatedAnnealing` (geometric cooling), all deterministic
under a seed and all honest about an exhausted budget.
- Programming: `MinimiseLinear` and `MinimiseLinearRows`, the revised
simplex with a two-phase start over standard-form and two-sided row
programs, and `MinimiseQP`, the active-set method for the strictly
convex program with the multipliers returned.
- Root finding: `FindRoot` (Brent), `FindRootBrent` for a scalar
bracketed root, `FindRootNewton` and `FindRootSystem` (damped
Newton with Armijo backtracking and an optional Broyden rank-one
update in place of repeated Jacobian builds). A solver that
exhausts its budget is refused with an error unless the
best-effort exit is requested by name.
**Automatic differentiation.**
- A reverse-mode graph over the arithmetic surface, the matrix
products (single and batched), the reductions, slicing,
concatenation, axis permutation and the Fourier transforms; every
float leaf accumulates through `Backward`.
- Complex tensors differentiate under the Wirtinger convention, the
loss stays real, and mixed real-complex graphs compose exactly
through the 2·Re narrowing.
- Second order: `Hessian` (forward-over-reverse) and
`HessianVectorProduct` in two gradient evaluations.
- On top of the graph: `MinimiseNewtonCG` (truncated-CG Newton with
an Armijo line search), `SampleHMC` (Hamiltonian Monte Carlo on any
differentiable unnormalised density) and `AdjointODE` (adjoint
sensitivities at the cost of one extra solve).
**Plotting.**
- The `plot` package: deterministic SVG line charts of computed
series. Linear axes with five ticks, one legend line per series, a
`Line` constructor straight from two rank-1 arrays through the
promotion ladder, and a byte-identical file on every run, so a
figure in a paper is compared exactly like any other computed
number.
**Distributed execution.**
- The experimental `spmd` package: explicit SPMD worlds, one program
on many ranks, over TCP between machines or in one process over
channels, launched, listened for and joined through `Launch`,
`Listen` and `Join`. The movement collectives `Broadcast`,
`Scatter`, `Gather` and `AllGather` move arrays between ranks; the
sharded reductions cut a global array on the canonical fold
partition's block boundaries and combine the partials through the
same balanced tree the single-array fold uses, so `Sum`, `Min`,
`Max`, `Any`, `All`, `Prod`, the norm and the dot families answer
the single-array reduction's exact bits at any world size, whatever
the order the frames arrive in; `Reduce` and `AllReduce` fold the
ranks' same-shaped arrays elementwise in rank index order.
`ExchangeHalos` and `ExchangeHalosOnGrid` hand each rank's boundary
slabs to the neighbours of a decomposition laid out on a row-major
process grid. Every failure or deadline fails the whole world
loudly, and no collective ever returns a partial numeric result.
**Data I/O.**
- CSV reading and writing, with or without a header row, every stored
numeric dtype written and read.
- FITS images with header cards in both directions, and binary and
ASCII table extensions.
- HDF5 in both directions: `LoadHDF5` reads the default and the
"latest" file formats (contiguous, compact and chunked storage, the
deflate, shuffle and fletcher32 filters, superblocks of versions 2
and 3 with the lookup3 checksum of each verified, group attributes
merged into each dataset), and `SaveHDF5` with `SaveHDF5Text`
writes every stored dtype at its native width, booleans through the
HDF5 enumeration convention, nested groups, attributes and optional
filters, byte-deterministic on every run. Unsupported format
features are refused by name, and cyclic or over-deep group walks
are refused.
- `LoadNetCDF`/`SaveNetCDF` for the NetCDF classic model (CDF-1 and
CDF-2), with named dimensions, text attributes and record
dimensions in both directions, and fixed-point variables landing at
their own width and sign.
- Memory mapping: `MapFloats`, `MapFloat32s` and `MapInts` open
native-endian files as read-only arrays without reading them, and
`SaveNativeFloats` writes the format they read.
**Examples.**
- Thirteen runnable workflows in `examples/`: ODE parameter fitting
by adjoint sensitivities, PSF deconvolution, HMC sampling, spectral
analysis, wavelet denoising, the exact pendulum period through
`EllipticK`, a Helmholtz system on the complex sparse solvers,
quasi-Monte Carlo integration, heat and wave evolution, regression
inference, a FITS star field, a NetCDF climate round trip and an
FFT tour.
**Project.**
- A determinism oracle pinning fixed workloads through the facade by
SHA-256 digest of the output bits, and a resource-leak harness
holding the goroutine count and live heap to baseline under
repeated heavy runs.
- Gitea Actions pipelines for test, race and release, with the
release notes extracted from this file's matching section.
- The document set: this changelog, the README, the API reference,
the architecture, the development guide, the benchmarking method,
the contribution rules and the security policy, under the MIT
licence.
+125
View File
@@ -0,0 +1,125 @@
# Contributing
Contributions to **Tensor** are governed by the Contributor terms
below; submitting one means you accept them.
## Contributor terms
1. This project belongs to its owner alone. The owner decides what is
accepted, in what form and when; the decision is final and needs no
justification.
2. By submitting a contribution you assign to Petr Balvín
<opensource@petrbalvin.org> all present and future copyright and
related rights in it, worldwide, for the full term of the rights,
with the right to relicense and sublicense without restriction,
including under proprietary terms.
3. Where that assignment is not effective, it counts as a perpetual,
irrevocable, royalty-free licence with the same scope.
4. To the fullest extent permitted by law, you waive any right of
attribution and integrity in the contribution. The project names no
contributors and keeps no credits list.
5. By submitting you represent that the work is yours and that you
hold the rights to assign it as above.
## Development setup
Requirements: Go 1.27.1 or newer (the version `go.mod` pins), and
[just](https://github.com/casey/just) for the recipes. A C compiler is
needed for the race detector, which `just test`'s sibling `just race`
and `just gates` run.
```sh
git clone https://sourcedock.dev/petrbalvin/tensor.git
cd tensor
just build
just test
```
The module has zero third-party dependencies: the standard library
covers everything, and a new dependency needs a reason that survives
review.
## Workflow
1. Branch from `development`. Never commit directly to `main`, which is release-only.
2. Commit in [Conventional Commits](https://www.conventionalcommits.org/) form:
`type(scope): description`, subject line only, imperative mood, lowercase after the
colon, no trailing full stop. Allowed types: `feat`, `fix`, `docs`, `style`,
`refactor`, `perf`, `test`, `chore`, `ci`, `build`, `revert`.
3. One logical change per commit. A refactor, a behaviour change and a formatting pass
are three commits, never one.
4. Record every user-visible change in `CHANGELOG.md` under `## [development]`.
5. Add or update tests. Coverage stays at 80 percent or more; it is a hard gate.
6. Update the documentation when the public API, the configuration or the behaviour
changes. The reader-facing reference is `docs/API.md`, one section per package, and
a new exported symbol belongs in it.
7. Open a pull request against `development`.
Releases are cut by merging `development` into `main` and tagging `vX.Y.Z`. The release
workflow runs the gates at the tag and publishes the release with its notes.
## Code style
`gofmt` and `go vet` run through `just fmt` and `just vet`, with zero diff and zero
warnings tolerated. `just gates` is the definition of done in one command, and the recipe
file names what it contains. Errors are checked explicitly, wrapped with `%w` so the
cause stays inspectable, prefixed `tensor: ` so the origin is always the library's, and
nothing panics outside `main`. The `golang` skill holds the rules the project follows;
the recipe file holds the commands.
New source files open with the project's two-line licence header, whose SPDX
identifier matches `LICENSE`. Configuration files, workflows and dotfiles do not carry
it.
## AI contribution policy
AI tools are welcome as productivity aids and are a normal part of modern software
development. What matters is that the contribution stays understandable, reviewable and
genuinely useful.
- **Disclose the assistance.** If AI helped draft any part of a commit, issue, pull
request or review, say so.
- **Commit messages carry exactly one trailer**, on the line after the subject:
```
Assisted-by: GLM 5.3
```
Name the model that did the work, spelled the way its maker spells it, for example
`GLM 5.3`, `DeepSeek V4.1 Flash` or `Qwen 3.8 Flash`. No `Co-Authored-By`, no `Signed-off-by`,
no other trailers, and no prose: the trailer is the disclosure.
- **Issues and pull requests** attribute the assistance in a comment, for example
`_Assisted-by: GLM 5.3_`. It does not belong in the pull request description.
- **Take responsibility.** You are accountable for the accuracy, completeness and
intent of everything you submit, whether or not AI produced it.
- **Review before marking ready.** Read the diff carefully, run it locally, and add the
tests it needs. Do not mark a pull request ready until you can defend every change in
it.
- **Quality over quantity.** Contributions that look like un-reviewed output, or whose
author cannot engage substantively during review, may be closed.
- **Preferred models.** Prefer open-weight models with transparent training data and
minimal output filtering.
AI assists. It does not replace judgement.
## Continuous integration
Workflows live in `.gitea/workflows/` and run on the project's own runners:
| Workflow | Trigger | What it does |
|---|---|---|
| Test | push or pull request to `development` | build, format check, vet, the test suite with the coverage floor, the oracle digests and a benchmark smoke run |
| Race | dispatched by hand | the suite under the race detector, with the oracle digests across glibc and musl |
| Release | a `v*` tag | the same gates minus race at the tag, then the release with its notes from `CHANGELOG.md` |
The local equivalent is `just gates`, which is the same set plus the race detector in
one command.
## Reporting bugs
Open an issue at `https://sourcedock.dev/petrbalvin/tensor/issues` with the
version, the operating system and architecture, the exact command, the full output,
and the expected against the actual behaviour.
**Security issues do not go in the issue tracker.** Report them as
[SECURITY.md](SECURITY.md) describes.
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+727
View File
@@ -0,0 +1,727 @@
# Tensor
Tensor is a scientific computing library for Go: n-dimensional arrays,
dense and sparse linear algebra, differential equation solvers,
quadrature, statistics, signal transforms, optimisation, deterministic
scientific charts and a reverse-mode differentiable core, built for the
natural sciences: cosmology, astronomy, quantum, particle and nuclear
physics, condensed matter and materials science, chemistry, biology and
genetics. It is pure Go with no cgo, no GPU stack and no third-party
dependency, and it is deterministic by contract: parallel kernels
reduce in a fixed order, so the same program with the same seed
produces bit-identical output run to run.
## What Tensor optimises for
Tensor is not built to win a speed contest. It keeps no benchmark
tables against other libraries, NumPy included: they are not the
competition, and racing them would settle nothing. The only
performance numbers in this repository compare Tensor against its own
previous revisions, under the discipline in
[docs/BENCHMARKING.md](docs/BENCHMARKING.md), so that a change is
judged by what it costs and never by an impression. Speed is an
engineering duty, not the goal.
The goal is a tool a scientist can trust with their numbers:
- **Determinism, everywhere, without exceptions.** The same program
with the same seed produces bit-identical output, run to run, on any
core count and under any worker setting. Parallel kernels reduce in
an order derived from the data alone, never from the machine, and
the generator is pinned and stable across Go releases. A number that
moves between runs is not a result.
- **Portability, without variants.** Pure Go on the standard library:
no cgo, no third-party dependency. One portable build, no build
flags, no CPU feature probes; there is no fast edition and no slow
edition of the truth to keep in agreement.
- **Precision, measured rather than claimed.** Narrow element types
accumulate in float64, and where two algorithms round differently
the choice is settled against high-precision referents, `big.Float`
arithmetic at hundreds of bits. Every solver carries its
exact-reference or residual check in the test suite.
This is also why the GPU is not the path for Tensor. GPU execution is
non-deterministic by construction: the order of a reduction follows
the hardware topology and the device scheduler instead of the data,
and the arithmetic carries no exactness guarantee to hold, the
graphics APIs themselves permitting results several ULP off. A
library contracted to exact reproducibility does not hand its numbers
to hardware that cannot sign for them.
## Features
- **Arrays**: immutable, row-major, n-dimensional arrays of int64,
IEEE 754 half precision, float32, float64 and complex128 with a
strict promotion ladder and loud shape errors; slicing views,
gathers, scatters, sorting, Einsum, interpolation and the special
functions from the gamma family to the elliptic integrals.
- **Linear algebra (`linalg`)**: dense factorisations and
eigenproblems in real, complex and general arithmetic, matrix
functions, regularised and truncated solves, the rank-revealing QR,
sparse CSR solvers with ILU preconditioning, sparse direct
Cholesky and LU with fill-reducing orderings and the rank-one
update and downdate, LSQR and LSMR least squares, Lanczos and
Arnoldi eigensolvers, the Krylov action of a matrix exponential,
polynomial fitting and natural cubic splines.
- **Differential equations (`integrate`)**: adaptive Dormand-Prince,
stiff systems from variable-step BDF2 up to variable-order BDF 1 to
5, the L-stable Rosenbrock-Wanner ROS4, index-1
differential-algebraic systems in mass-matrix form, symplectic
integrators from velocity Verlet through Yoshida's fourth order to
the implicit midpoint, event detection, boundary values by shooting
and by adaptive Lobatto collocation, Gauss-Legendre quadrature,
globally adaptive cubature, heat and wave evolution in one and two
space dimensions, flux-limited advection, and finite-element
Poisson solves on triangular and tetrahedral meshes.
- **Signal and transforms (`signal`)**: FFTs of any length,
multi-dimensional and real-input transforms, cosine and sine
transforms, the NUFFT, Welch, spectrogram and Lomb-Scargle spectra,
the Hilbert envelope, decimation and resampling, Butterworth,
Chebyshev, inverse Chebyshev and elliptic filter design, the window
catalogue, zero-phase filtfilt, median and rank filters, Haar and
Daubechies wavelets, continuous wavelets, Kalman filtering, AR and
ARMA estimation, convolutions, pooling and stencils, and spectral
Poisson solves.
- **Statistics (`stats`)**: CDFs, quantiles and draws for the normal,
exponential, gamma, chi-square, Student t, Poisson and binomial
laws; density, CDF and quantile for Weibull, lognormal and Pareto;
the negative binomial PMF, CDF and quantile; the Dirichlet density,
mean, mode and draws; the noncentral chi-square, F and t families
through their density, CDF and quantile; histograms and rolling
windows, robust descriptives, Welch's t-test, Kolmogorov-Smirnov,
Mann-Whitney U, one-way ANOVA, bootstrap intervals, rank
correlations and multiple-testing corrections, multivariate
normals, kernel density, Gaussian processes, clustering, PCA, and
linear, weighted, logistic, Poisson, lasso, elastic-net, Huber and
quantile regression with the classical inference beside Theil-Sen's
robust pair.
- **Optimisation (`optim`)**: Levenberg-Marquardt with an optional
analytic Jacobian, L-BFGS with box bounds, linearly constrained
minimisation through the augmented Lagrangian with nonlinear
equality and inequality rows, Nelder-Mead, differential evolution,
CMA-ES and simulated annealing, the revised simplex and active-set
quadratic programming, and Brent, Newton, Broyden quasi-Newton and
damped-Newton root finding.
- **Differentiable core (`grad`)**: a reverse-mode graph over the
arithmetic surface, the matrix products and the transforms, complex
Wirtinger differentiation, Hessians, Newton-CG on the graph,
Hamiltonian Monte Carlo and adjoint ODE sensitivities.
- **Data (`io`)**: CSV, FITS images and tables, HDF5 datasets read
and written (both superblock generations, chunked storage, deflate
and shuffle filters, attributes and string data), the NetCDF
classic model, and memory-mapped arrays that let a data cube far
larger than RAM open instantly.
- **Worked examples**: thirteen runnable workflows in `examples/`,
each solving its problem end to end: ODE parameter fitting by
adjoint sensitivities, PSF deconvolution, Hamiltonian Monte Carlo
sampling, spectral analysis, wavelet denoising, the exact pendulum
period through the elliptic integral, a Helmholtz system on the
complex sparse solvers, quasi-Monte Carlo integration, heat and
wave evolution, regression inference, a FITS star field, a NetCDF
climate round trip, and an FFT tour.
## Experimental distributed computing
**This is an experiment.** The `spmd` package carries a revolutionary
technological concept: one program running on many ranks, in one
process or across machines, where the order of a reduction is a
function of the data alone, so a sharded reduction carries the
single-array reduction's exact bits at any world size. The concept is
not fully verified and remains the subject of research. The
single-machine library is the settled product; the distributed surface
is new, its behaviour and performance on real networks are not yet
measured, and it is expected to change as the research moves.
What the experiment carries today, and what already holds: the
movement collectives (`Broadcast`, `Scatter`, `Gather`, `AllGather`)
and two families of reductions whose answers the test suite, the
oracle digests and the loopback cluster tests in CI pin against the
single-array reductions they must equal. What the suite cannot reach,
the bandwidth and behaviour of real networks above all, is exactly
what the research is for. Read the contract in
[docs/API.md](docs/API.md) before you build on it, and treat its edges
as open questions rather than finished answers.
The shortest form of the experiment, four ranks over one million
values, with the sharded sum equal to the single-array sum bit for
bit:
```go
package main
import (
"fmt"
"sourcedock.dev/petrbalvin/tensor"
"sourcedock.dev/petrbalvin/tensor/spmd"
)
func main() {
err := spmd.Launch(4, func(w *spmd.World) error {
const n = 1000000
vals := make([]float64, n)
for i := range vals {
vals[i] = float64((i*7919)%2001-1000) / 7.0
}
whole, _ := tensor.FromFloats(vals, n)
span, err := spmd.Partition(n, w.Size(), w.Rank())
if err != nil {
return err
}
local, err := tensor.Slice(whole, 0, span.Lo, span.Hi)
if err != nil {
return err
}
got, err := w.AllReduceShards(local, span, spmd.Sum)
if err != nil {
return err
}
if w.Rank() == 0 {
want := tensor.Sum(whole)
fmt.Printf("size %d: sharded %v equals single-array %v: %v\n",
w.Size(), got.Float(), want.Float(), got.Float() == want.Float())
}
return w.Barrier()
})
if err != nil {
panic(err)
}
}
```
`Launch` swaps for `Listen` and `Join` when the ranks are processes on
different machines, and nothing else in the program moves.
## Install
As a library:
```sh
go get sourcedock.dev/petrbalvin/tensor
```
Requires Go 1.27.1 or newer, the exact version `go.mod` declares.
One import covers everything: the root package re-exports the exported
surface of every domain package, so `tensor.SVD` and `linalg.SVD` name
the same function. A domain package may also be imported on its own
for a narrow dependency graph; the general array constructors live in
the root package (`linalg.ArrayFromFloatsSafe` is the one exported
outside it), and the arrays are the same type either way, because
`tensor.Array` is an alias for the core array, not a wrapper around
it. Nothing is lost by mixing the two styles.
## Quick start
A stiff relaxation with a slow forcing, integrated to the analytic
answer:
```go
package main
import (
"fmt"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
// y' = -1e5·(y - cos t): a transient of width 1e-5 under a slow
// forcing. The adaptive BDF2 steps over the transient and then
// follows the forcing; an explicit scheme is pinned to the
// stability limit h < 2e-5 for the whole run.
f := func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{-1e5 * (y.FloatAt(0) - math.Cos(t))}, 1)
}
y0, _ := tensor.FromFloats([]float64{0}, 1)
end, err := tensor.IntegrateBDF2(f, 0, 1, y0, tensor.ODEOptions{MaxSteps: 2000})
if err != nil {
panic(err)
}
exact := (1e10*math.Cos(1) + 1e5*math.Sin(1)) / (1e10 + 1)
fmt.Printf("y(1) = %.12f, error %.2e\n", end.FloatAt(0),
math.Abs(end.FloatAt(0)-exact))
}
```
The stiff solver lands on the exact value `(k²·cos 1 + k·sin 1)/(k² + 1)`
inside a 2000-step budget where the explicit pair, pinned to its
stability limit, cannot follow the run at all. It prints
`y(1) = 0.540310721007, error 4.83e-10`.
## Usage
Tensor's behaviour is a contract, not a convention:
- **Immutable arrays.** No operation mutates its inputs; results are
fresh arrays, and views never alias a buffer a later step could
rewrite. A returned array is yours alone, which is also what makes
them safe to share between goroutines.
- **Errors, not lies.** Singular systems, exhausted step budgets,
malformed files and impossible shapes come back as errors prefixed
`tensor: ` that say what happened. Nothing is silently truncated,
clamped or filled.
- **Determinism.** The generator is xoshiro256++ seeded through
splitmix64, stable across Go releases, because the standard library
does not promise stable output and reproducibility is the point of
a seed. Parallel kernels keep their reduction order fixed, so a
result does not move with the core count; `SetNumCPU(n)` pins the
worker count for containers and small machines.
- **Scientific scope.** Tensor is for the natural sciences and for
nothing else. It carries nothing for artificial intelligence,
economics or finance: no neural-network machinery, no training
loops, no market or portfolio helpers.
The convolution and pooling functions are signal-processing
stencils (PSF deconvolution, image filtering), not model layers.
Differentiation exists because fitting parameters to data and
sensitivity analysis are scientific tools.
What follows is one short program per package, each complete and
runnable as written. The full surface, option by option, is in
[docs/API.md](docs/API.md), and the longer workflows are the thirteen
programs in [`examples/`](examples).
### Arrays
```go
package main
import (
"fmt"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
y, err := tensor.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
if err != nil {
panic(err)
}
// A slice is a read-only view: no copy, and no way for a later
// step to write through it.
cols, _ := tensor.Slice(y, 1, 1, 3)
// Reductions name the axis, and the axis disappears from the shape.
rowSum, _ := tensor.SumAxis(y, 1)
// The dtype ladder promotes on request, never implicitly downward.
halves, _ := tensor.Astype(y, tensor.Float32)
fmt.Println(y.Shape(), rowSum, halves.Dtype())
fmt.Println(cols)
fmt.Println(y)
}
```
`Shape` reports the extents, `Dtype` the element type, `Len` the
element count. Element-wise operations (`Add`, `Mul`, `Exp`, `Sqrt`,
the comparisons, `Where`) and the shape moves (`Reshape`, `Transpose`,
`Concat`, `Stack`, `Pad`) all take and return whole arrays, so a
formula reads as one expression per line.
### Linear algebra
```go
package main
import (
"fmt"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
a, _ := tensor.FromFloats([]float64{4, 1, 1, 3}, 2, 2)
b, _ := tensor.FromFloats([]float64{1, 2}, 2)
// One LU with partial pivoting; a singular matrix is an error.
x, err := tensor.Solve(a, b)
if err != nil {
panic(err)
}
// The symmetric eigenproblem returns ascending values and the
// orthonormal eigenvectors as columns.
values, vectors, _ := tensor.Eigen(a)
// The Cholesky factor for reuse across right-hand sides.
l, _ := tensor.Cholesky(a)
fmt.Println(x, values)
fmt.Println(l, vectors.Shape())
}
```
Sparse systems go through the same array type: `SparseFrom` builds the
COO form, `CSRFromCOO` and `CSCFromCOO` the compressed views,
`NewSparseCholesky` and `NewSparseLU` the direct factorisations under a
fill-reducing ordering, `NewSparseILU` the preconditioner, and
`SpSolve`, `SpSolveBiCGSTAB`, `SpLSQR` and `SpEigen` the iterative
solvers.
### Differential equations
```go
package main
import (
"fmt"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
// The harmonic oscillator as a first-order system: y = (position,
// velocity), so y' = (velocity, -position).
f := func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
y0, _ := tensor.FromFloats([]float64{0, 1}, 2)
// Event detection is a watch on the trajectory, so the crossing
// time comes from the interpolant rather than from the step grid.
hits, end, err := tensor.IntegrateODEEvents(f, 0, 4, y0,
[]tensor.ODEWatch{{
Function: func(t float64, y *tensor.Array) (float64, error) {
return y.FloatAt(0), nil
},
Direction: -1, // falling crossings only
}}, tensor.ODEOptions{})
if err != nil {
panic(err)
}
fmt.Printf("first minimum at t = %.6f, want pi = %.6f\n", hits[0].Time, math.Pi)
fmt.Printf("y(4) = %.6f\n", end.FloatAt(0))
}
```
`IntegrateODE` is the adaptive Dormand-Prince pair, `IntegrateBDF2`
and `IntegrateBDFVar` the stiff routes, `IntegrateROS4` the L-stable
one, `IntegrateDAE` the mass-matrix form. Quadrature is
`IntegrateFunction` and `IntegrateND`, the boundary value problems are
`IntegrateBoundary` and `SolveBoundaryCollocation`, and the turnkey
time steppers are the `IntegrateHeat1D/2D`, `IntegrateWave1D/2D` and
advection families.
### Transforms, filters and spectra
```go
package main
import (
"fmt"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
fs = 1000.0
n = 1000
)
samples := make([]float64, n)
for i := range n {
t := float64(i) / fs
samples[i] = math.Sin(2*math.Pi*50*t) + 0.25*math.Sin(2*math.Pi*120*t)
}
x, _ := tensor.FromFloats(samples, n)
// Welch's averaged periodogram: windowed segments, one-sided.
freqs, psd, err := tensor.WelchPSD(x, fs, 256, 128, "hann")
if err != nil {
panic(err)
}
// Bin 0 is the mean; the peak above it is the tone.
rest, _ := tensor.Slice(psd, 0, 1, psd.Len())
idx, _ := tensor.ArgMax(rest)
fmt.Printf("peak at %.1f Hz, bin spacing %.1f Hz\n",
freqs.FloatAt(idx+1), freqs.FloatAt(1))
// A Butterworth design plus filtfilt: zero phase, so no lag.
b, a, _ := tensor.ButterworthLowPass(4, fs, 60)
clean, _ := tensor.Filtfilt(b, a, x)
peakIn, _ := tensor.Max(x)
peakOut, _ := tensor.Max(clean)
fmt.Printf("input peak %.3f, filtered peak %.3f\n", peakIn.Float(), peakOut.Float())
}
```
The Fourier family covers `FFT`/`IFFT` of any length, `FFT2`, `FFT3`,
`FFTN`, the real-input `RFFT`/`IRFFT`, the cosine and sine transforms,
the Haar and Daubechies wavelets, the analytic `CWT`, and the STFT,
spectrogram and Lomb-Scargle estimators beside Welch.
### Statistics
```go
package main
import (
"fmt"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
// y = 2 + 3x with noise, the intercept column first.
design, _ := tensor.FromFloats([]float64{
1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5,
}, 6, 2)
y, _ := tensor.FromFloats([]float64{2.1, 4.9, 8.2, 11.1, 13.8, 17.2}, 6)
fit, err := tensor.LinearRegression(design, y)
if err != nil {
panic(err)
}
fmt.Printf("slope %.3f ± %.3f, t = %.2f, p = %.2g, R2 = %.4f\n",
fit.Coefficients[1], fit.StandardErrors[1],
fit.TStatistics[1], fit.PValues[1], fit.RSquared)
// Distributions answer one call each, with the tail the caller asks for.
fmt.Printf("P(Z <= 1.96) = %.4f, t(10) 97.5%% = %.4f\n",
tensor.NormalCDF(1.96), mustQuantile(tensor.StudentTQuantile(0.975, 10)))
}
// Quantiles can fail on a parameter outside their domain, so the
// example carries the error rather than dropping it.
func mustQuantile(q float64, err error) float64 {
if err != nil {
panic(err)
}
return q
}
```
The models are `LinearRegression`, `WeightedLinearRegression`,
`LogisticRegression`, `PoissonRegression`, `Lasso`, `ElasticNet`,
`HuberRegression` and `QuantileRegression`, each returning its
coefficient table with the standard errors and p-values beside it;
`TheilSenRegression` returns the robust intercept and slope as a pair.
The tests are `WelchTTest`, `KolmogorovSmirnovTest`, `MannWhitneyU`,
`ANOVAOneWay` and `ChiSquareGoodnessOfFit`; the multivariate tools are
`PCA`, `KMeans`, `GaussianMixture` and `GaussianProcessRegression`.
### Optimisation
```go
package main
import (
"fmt"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
// Fit A·exp(−k·t) to six noisy observations by least squares.
ts := []float64{0, 1, 2, 3, 4, 5}
obs := []float64{2.0, 1.22, 0.74, 0.45, 0.27, 0.17}
residual := func(p *tensor.Array) (*tensor.Array, error) {
out := make([]float64, len(ts))
for i, t := range ts {
out[i] = p.FloatAt(0)*math.Exp(-p.FloatAt(1)*t) - obs[i]
}
return tensor.FromFloats(out, len(out))
}
p0, _ := tensor.FromFloats([]float64{1, 0.5}, 2)
p, ss, err := tensor.LevenbergMarquardt(residual, p0, tensor.LMOptions{})
if err != nil {
panic(err)
}
fmt.Printf("A = %.4f, k = %.4f, sum of squares %.3e\n",
p.FloatAt(0), p.FloatAt(1), ss)
}
```
`MinimiseLBFGS` with optional box walls, `Minimise` (Nelder-Mead),
`MinimiseConstrained` for the linear rows, `MinimiseNonlinearConstrained`
for constraint functions, `MinimiseDifferentialEvolution`,
`MinimiseCMAES` and `MinimiseSimulatedAnnealing` for the multimodal
landscapes, `MinimiseLinear` and `MinimiseQP` for the programs, and
`FindRoot`, `FindRootNewton` and `FindRootSystem` for the roots. A
solver that runs out of budget refuses rather than reporting a
converged answer, and `AllowBudgetExit` opts into the best point
instead; `MinimiseDifferentialEvolution` is the exception, returning
its best generation without an error.
### Data in and out
```go
package main
import (
"fmt"
"os"
"path/filepath"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
dir, err := os.MkdirTemp("", "tensor-readme")
if err != nil {
panic(err)
}
defer os.RemoveAll(dir)
scan, _ := tensor.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
path := filepath.Join(dir, "field.h5")
err = tensor.SaveHDF5(path, []tensor.HDF5Dataset{{
Path: "/scan/temperature",
Values: scan,
Attrs: map[string]string{"units": "K"},
}}, nil, tensor.HDF5WriteOptions{Gzip: 6, Shuffle: true})
if err != nil {
panic(err)
}
sets, err := tensor.LoadHDF5(path)
if err != nil {
panic(err)
}
fmt.Println(sets[0].Path, sets[0].Shape, sets[0].Attrs["units"], sets[0].Values)
}
```
`LoadCSV`/`SaveCSV` handle the tabular case, `LoadFITS`/`SaveFITS`
images and `LoadFITSTable`/`SaveFITSTable` the table extensions,
`LoadNetCDF`/`SaveNetCDF` the classic model, and `MapFloats`,
`MapFloat32s` and `MapInts` open a native-endian file as a read-only
array without reading it; `SaveNativeFloats` writes the float64 form
`MapFloats` reads.
### Charts
```go
package main
import (
"fmt"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
wavelengths := make([]float64, 200)
intensity := make([]float64, 200)
for i := range wavelengths {
wavelengths[i] = 400 + 2*float64(i)
intensity[i] = 100 + 40*math.Exp(-math.Pow(wavelengths[i]-589, 2)/25)
}
xs, _ := tensor.FromFloats(wavelengths, 200)
ys, _ := tensor.FromFloats(intensity, 200)
series, _ := tensor.Line("sodium D line", xs, ys)
chart := tensor.Chart{
Title: "Absorption spectrum",
XLabel: "wavelength [nm]", YLabel: "intensity",
Series: []tensor.Series{series},
}
if err := chart.WriteSVG("spectrum.svg"); err != nil {
panic(err)
}
fmt.Println("spectrum.svg written")
}
```
Deterministic SVG line charts: linear axes, five ticks each, one legend
line per series, and a byte-identical file on every run, so a figure in
a paper is compared exactly like any other computed number. The package
is small by intent; it draws the figures, it does not stage a cinema.
### Automatic differentiation
```go
package main
import (
"fmt"
"sourcedock.dev/petrbalvin/tensor/grad"
)
func main() {
x, _ := grad.FromFloat64s([]float64{1, 2, 3}, true, 3)
w, _ := grad.FromFloat64s([]float64{0.5, -1, 2}, true, 3)
prod, _ := x.Mul(w)
sq, _ := prod.Pow(2)
loss, _ := sq.Sum()
if err := loss.Backward(); err != nil {
panic(err)
}
// dL/dx = 2·x·w² and dL/dw = 2·x²·w, exact to rounding.
fmt.Println(x.Grad(), w.Grad())
// Second-order questions come from the same graph: H·v in two
// gradient evaluations, the tool Newton-CG scales on. The Hessian
// of Σz² is 2·I, so H·x is 2·x.
quadratic := func(z *grad.Tensor) (*grad.Tensor, error) {
sq, err := z.Pow(2)
if err != nil {
return nil, err
}
return sq.Sum()
}
hv, err := grad.HessianVectorProduct(quadratic, x, x, grad.HessianOptions{})
if err != nil {
panic(err)
}
fmt.Println(hv)
}
```
`MinimiseNewtonCG` minimises a graph function with truncated-CG Newton
steps, `SampleHMC` runs Hamiltonian Monte Carlo on any differentiable
unnormalised density, and `AdjointODE` differentiates an ODE solution
at the cost of one extra solve. Complex graphs follow the Wirtinger
convention, with `Real`, `Imag`, `Conj` and `Abs2` bridging into a real
loss.
at the cost of one extra solve. Complex graphs follow the Wirtinger
convention, with `Real`, `Imag`, `Conj` and `Abs2` bridging into a real
loss.
### Where to look next
- The complete surface, package by package: [docs/API.md](docs/API.md).
- Runnable end-to-end workflows, one directory each:
[`examples/`](examples).
- The package map and the data flow:
[docs/ARCHITECTURE.md](docs/ARCHITECTURE.md).
- Building, testing, benchmarks and releases:
[docs/DEVELOPMENT.md](docs/DEVELOPMENT.md) and
[docs/BENCHMARKING.md](docs/BENCHMARKING.md).
- The godoc comments in the source are the authority on signatures:
`go doc -all sourcedock.dev/petrbalvin/tensor`.
## Development
```sh
just build # compile everything, examples included
just test # the suite with the coverage floor
just gates # the definition of done: build, fmt-check, vet, test, race
```
CI (Gitea Actions) enforces the same gates on every push to
`development`, race excepted: the race detector is dispatched by hand.
See [docs/DEVELOPMENT.md](docs/DEVELOPMENT.md) for the full workflow
and [CONTRIBUTING.md](CONTRIBUTING.md) for how to contribute.
## Documentation
- [docs/API.md](docs/API.md): the exported API reference, per package
- [docs/ARCHITECTURE.md](docs/ARCHITECTURE.md): the package map,
data flow and design
- [docs/DEVELOPMENT.md](docs/DEVELOPMENT.md): building, testing and
releasing
- [docs/BENCHMARKING.md](docs/BENCHMARKING.md): how performance is
measured, and the reports
- [CHANGELOG.md](CHANGELOG.md): release history
## Licence
MIT. See [LICENSE](LICENSE) for the text.
Copyright © 2026 [Petr Balvín](https://petrbalvin.org)
+39
View File
@@ -0,0 +1,39 @@
# Security policy
## Supported versions
Security fixes go to the newest release and to the `development` branch. Older releases
do not receive them.
| Version | Supported |
|---|---|
| 1.0.0 | yes |
| older releases | no |
## Reporting a vulnerability
**Do not open a public issue for a security problem.** A public report tells everyone
about the flaw before there is a fix. Report it privately to
**opensource@petrbalvin.org**.
Include:
- the version or commit you tested, and the platform
- what the problem is, and what an attacker gains from it
- the smallest reproducer you have, ideally a test or a single command
- a suggested fix, if you have one
## What to expect
- A human reads the report, and you get an acknowledgement.
- You are kept informed while the fix is being made, and told when it ships.
- The fix is released before the details are published, and the timing is agreed with
you.
- Thanks are given by private acknowledgement: the project names no contributors in
its published history.
## Out of scope
- Findings that require the attacker to already run code as the user, or to have local
access.
- Missing hardening with no demonstrated impact.
+58
View File
@@ -0,0 +1,58 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package tensor is the public facade of the library: every exported
// symbol of every domain package is re-exported here under one name,
// so user code imports only "sourcedock.dev/petrbalvin/tensor" and
// writes tensor.Everything. The implementations live in the domain
// packages (linalg, signal, integrate, stats, optim, io, grad) and the
// shared array core in internal/core.
//
// # What it is
//
// A scientific computing library in pure Go: no cgo, no GPU stack, no
// third-party dependencies. Arrays over int64, float32, float64 and
// complex128 with a strict promotion ladder and loud shape errors;
// dense and sparse linear algebra through eigensolvers, decompositions
// and Krylov methods; Fourier, cosine/sine, wavelet and continuous
// wavelet transforms; ordinary differential equations with events and
// adjoint sensitivities; PDE evolution in one and two space dimensions;
// quadrature and cubature; distributions, inference and linear
// regression with classical errors; global and local optimisation; and
// a reverse-mode differentiable core that covers the arithmetic, the
// transforms and the second-order questions alike.
//
// # A tour in one breath
//
// z, _ := tensor.MatMul2D(a, b) // dense linear algebra
// spec, _ := tensor.FFT(x) // transforms, any length
// end, _ := tensor.IntegrateODE(f, 0, 1, y0, tensor.ODEOptions{})
// res, _ := tensor.LinearRegression(X, y) // full inference
// loss.Backward() // exact gradients
// xs, _ := tensor.MinimiseNewtonCG(f, x0, tensor.NewtonCGOptions{})
//
// The examples in this documentation are executable and checked by
// the test suite; the examples/ directory carries the longer
// workflows (ODE parameter fitting, PSF deconvolution, HMC sampling,
// spectral analysis).
//
// # The guarantees
//
// - Deterministic: parallel kernels reduce in a fixed order, so a
// given element order gives bit-identical results run to run;
// SetNumCPU pins the parallelism.
// - Loud: a shape mismatch, a singular matrix, an exhausted solver
// budget or a CFL violation is an error naming itself, never a
// silently wrong number.
// - Immutable: operations never modify their inputs.
// - Reproducible: the generator is xoshiro256++ seeded through
// splitmix64, stable across Go releases, and Substream hands out
// the provably distinct members of one seed's stream family.
//
// # Conventions
//
// Functions return (value, error) and wrap errors with context;
// scalars come back as Scalar when the dtype follows the input. The
// names are the library's own: consistent with the established
// patterns here rather than borrowed from any other array library.
package tensor
+3583
View File
File diff suppressed because it is too large Load Diff
+207
View File
@@ -0,0 +1,207 @@
# Architecture
How Tensor is put together. Every node, package and arrow below
exists in the source tree; nothing is aspirational.
## Overview
```mermaid
flowchart TD
X["examples, thirteen main programs"]
F["tensor, the root facade"]
G["grad"]
I["integrate"]
L["linalg"]
S["signal"]
ST["stats"]
O["optim"]
IO["io"]
C["internal/core"]
B["internal/base"]
E["internal/engine"]
X --> F
F --> G
F --> I
F --> L
F --> S
F --> ST
F --> O
F --> IO
G --> S
G --> I
I --> L
I --> O
O --> L
G --> C
I --> C
L --> C
S --> C
ST --> C
O --> C
IO --> C
C --> B
C --> E
B --> E
```
Tensor is a scientific computing library in pure Go: no cgo, no GPU
stack, no third-party dependencies. The module splits into one core
package and one package per domain, with a strict dependency
direction: domains depend on the core, never the other way round, and
nothing below the root imports the root.
The root package is a facade, and its mechanism is a type alias plus a
forward: `type Array = core.Array` makes the root's array and the
core's one the same type rather than a wrapper, and
`facade_generated.go` declares the rest as `var SVD = linalg.SVD` and
its neighbours, so `tensor.SVD` and `linalg.SVD` are one function
value. A new domain export needs a line in that file. Nothing else in
the root package holds logic.
Every arrow above is one import. Two groups are elided for
readability: each domain also imports `internal/base` beside
`internal/core`, and the three domains that drive the parallel
fan-out themselves (`linalg`, `signal`, `grad`) also import
`internal/engine` directly.
## Packages
| Package | Responsibility |
|---|---|
| `tensor` (root) | the facade: re-exports every domain symbol through `facade_generated.go` and owns no logic beyond the alias definitions |
| `internal/core` | the `Array` type and everything that treats it as an n-dimensional value: constructors, element-wise arithmetic, reductions, shape moves, indexing and views, sorting, Einsum, interpolation, sparse COO, special functions, quasirandom sequences, the reproducible generator; deliberately no domain knowledge |
| `internal/base` | the shared primitives: the generic LU (`Factor`, `SolveSystem`), error construction with the `tensor: ` prefix, shape formatting, machine epsilon, so no domain imports another for plumbing; deliberately no array knowledge |
| `internal/engine` | the parallel scheduler and the pooled scratch buffers every kernel fans out through; the only place that starts workers |
| `linalg` | dense and sparse linear algebra: factorisations, eigenproblems, matrix functions, regularised and truncated solves, iterative sparse solvers and eigensolvers, polynomial fitting and cubic splines |
| `signal` | transforms and stencils: Fourier, cosine and sine transforms, the NUFFT, spectral estimation, filter design, wavelets, convolutions, pooling, and the spectral Poisson solves |
| `integrate` | differential equations and quadrature: ODE steppers with events, boundary-value shooting, Gauss rules, cubature, turnkey heat and wave evolution, finite elements |
| `stats` | distributions and inference: CDFs, quantiles, draws, descriptives, tests, regression models, multivariate normals, kernel density |
| `optim` | fitting and root finding: local, bounded, constrained and global minimisation |
| `plot` | deterministic SVG line charts of computed series: linear axes, legends, byte-identical figures |
| `spmd` | explicit SPMD worlds over TCP or in process: rank-0 routed links, the movement collectives and the sharded reductions whose answers are the single-array fold's exact bits at any world size; imported explicitly, the root facade does not re-export it |
| `io` | data formats: CSV, FITS images and tables, HDF5 datasets, the NetCDF classic model, memory-mapped arrays |
| `grad` | the reverse-mode differentiable core over the shared surface, second-order tools, Newton-CG, Hamiltonian Monte Carlo and adjoint sensitivities |
| `examples/*` | thirteen `main` packages, one workflow each, compiled by `just build`; they hold no library code and no tests |
The boundaries are as deliberate as the responsibilities. Five
domain-to-domain edges exist, each one-directional and each earning
its keep: `grad` reads `signal` for the spectral autograd nodes,
`grad` reads `integrate` for the adjoint ODE (imported under the name
`ode`), `integrate` reads `optim` for the shooting solver and the
implicit midpoint stage, `integrate` reads `linalg` for the
tridiagonal solve the Crank-Nicolson heat evolution rides and for the
sparse Cholesky and its orderings behind the FEM Poisson solver, and
`optim` reads `linalg` for `ArrayFromFloatsSafe`, the copying
constructor its callbacks are handed. The linear algebra `optim`
itself stands on is the generic LU in `internal/base`, not a domain
package. Beyond those edges, a domain never imports another domain,
the core never imports a domain, and a new cross-domain edge needs a
reason of the same kind. `spmd` is a domain in the dependency sense, it
reads `internal/core` and `internal/base` and nothing else of the
library, and it is deliberately absent from the root facade: a
distributed program imports it explicitly, the one place the
distributed surface is named. The canonical reduction partition and
the block folds it rides on live in `internal/core` beside the folds
themselves, so the single-array and the sharded reduction are one
computation by construction, not by a test.
The test-level harnesses sit at the root rather than in a package:
`oracle_test.go` pins a raw-bit digest per domain and platform,
`leak_test.go` measures the heap across repeated blocks, and
`example_test.go` carries the runnable godoc examples. Each domain
package carries its own `example_test.go` beside its tests, so the
documented call sequences are compiled and executed by `go test`.
## Data flow
```mermaid
flowchart TD
A1["constructors<br/>FromFloats, Zeros, Grid"]
A2["io loaders<br/>CSV, FITS, HDF5, NetCDF, mmap"]
C1["element-wise ops"]
C2["reductions"]
C3["domain kernels<br/>linalg, signal, integrate,<br/>stats, optim"]
C4["autograd graph<br/>grad.Backward"]
S1["spmd collectives<br/>Broadcast, Scatter, Gather,<br/>AllReduceShards"]
A1 --> C1
A1 --> C2
A1 --> C3
A1 --> C4
A1 --> S1
A2 --> C3
C3 --> C4
C1 --> C2
S1 --> C2
```
One kernel call, start to finish, as the layers see it:
```mermaid
sequenceDiagram
participant Caller
participant Facade as root facade
participant Domain as domain kernel
participant Core as internal/core
participant Eng as internal/engine
Caller->>Facade: tensor.SVD(a)
Facade->>Domain: linalg.SVD(a)
Domain->>Core: read payloads, allocate the output
Domain->>Eng: split the worker ranges
Eng-->>Domain: disjoint chunks, fixed order
Domain-->>Facade: fresh Array, never the input
Facade-->>Caller: the factorisation or an error
```
Arrays are dense and contiguous by construction: element i of an
array is payload index i. `Slice` preserves the invariant by a
rebased-pointer view where the selection keeps the trailing elements
contiguous (a slice along the leading axis, or one covering the whole
extent) and by a materialised copy everywhere else, and a strided
source is materialised before either path runs, because both assume
`payload[i]` is element i. A kernel that meets a non-contiguous input
through another route receives it materialised at the boundary, so
the audit has one rule: every kernel reads payloads assuming density.
That is what lets domain kernels read `RawFloats()` directly with no
per-element dispatch.
Views are read-only: nothing in the library writes through an array
it did not allocate, and optimiser updates route through
materialised parameters, so a view can never alias a buffer a later
step rewrites.
Errors are produced at the layer that detects them, prefixed
`tensor: ` by the shared error constructor in `internal/base`, and
returned unwrapped to the caller; no layer logs another layer's
error, swallows one, or turns one into a silently wrong number.
## State and lifetime
- **Long-lived.** The engine's worker pool, sized once by `SetNumCPU`
(the machine's core count by default) and re-pinnable at any time,
and the cached constant tables (Fourier twiddles, quadrature nodes)
held at package level with fixed contents.
- **Per-call.** Every kernel's output arrays and its index scratch;
nothing survives the call except the pool below.
- **Pooled.** Float64 scratch buffers flow through a typed pool whose
borrow path zeroes the window, closing the stale-buffer bug class
at the source; retention is capped, so one large table cannot pin
memory across the machine.
- **Concurrency.** Arrays are immutable and safe for concurrent use;
the `Generator` is not safe for concurrent use and is meant to be
owned by one goroutine. Parallel kernels keep their reduction order
fixed, so parallel results are bit-identical to serial ones, and
the determinism oracle at the root holds that contract by digest.
## Dependencies
There are no third-party dependencies: `go.mod` requires the standard
library alone, which is a property of the project, not an accident.
The two internal support packages exist to keep it that way and to
keep the dependency arrow one-directional: `internal/base` holds the
generic LU and the shared error and formatting helpers so no domain
imports another for plumbing, and `internal/engine` is the only place
that schedules workers, so every kernel's parallelism is decided in
one file. The five domain-to-domain imports named above are the whole
graph beyond that; a new one needs the same kind of reason.
+82
View File
@@ -0,0 +1,82 @@
# Benchmarking
How Tensor's performance is measured. The numbers a reader quotes must
be reproducible by following this document; anything else is an
impression, not a result.
## The tool
The benchmarks live beside the code they measure as Go benchmark
functions (`func BenchmarkXxx(b *testing.B)`), four hundred and
seventy-one of them across the packages: the array kernels and
parallel scheduler in `internal/core` and `internal/engine`, the dense
and sparse solvers in `linalg`, the transforms and filters in
`signal`, the differentiable core in `grad`, the integrators in
`integrate`, the optimisers in `optim`, the statistics in `stats`, the
collectives in `spmd` and the reader and writer round trips in `io`.
Run the whole suite with:
```sh
just bench
```
which runs `go test -run '^$' -bench=. -benchmem -count=5` over every
logic package. One package at a time:
```sh
go test ./linalg/ -bench 'BenchmarkSolve' -benchmem -count=5 -run xxx
```
`-benchmem` is not optional: allocations per operation are part of the
result. A kernel whose allocations grow has regressed even when its
time did not.
## The discipline
- **One process, A or B.** Two runs of two different binaries differ
by more than the effect being measured. When comparing a change
inside one revision, run both variants inside one process, or
interleave the sub-benchmarks behind a package-level switch.
- **Across revisions, interleave the rounds.** A release against the
head tree is necessarily two binaries; `just bench-report` runs the
two sides in alternating order over four rounds, so a host that
penalises the first run of a pair penalises both sides equally.
- **Idle machine.** A loaded machine profiles and times whatever ran
last. Close everything; treat any run sharing the box with other
work as void.
- **Five counts, median.** `just bench` takes five counts; report the
median and the spread. One to two percent is noise.
- **Deterministic inputs.** Every benchmark builds its inputs from the
seeded generator or fixed literals, so a number is tied to a
revision, not to a dice roll.
- **Complexity, not folklore.** A claim that a kernel is O(n log n)
belongs next to the measurements that show the scaling (two or
three sizes), not as an adjective.
## What is exact, what is fast
Performance numbers say nothing about correctness. The correctness
dossier lives in the test suite: `TestOracle` pins the raw-bit digest
of one fixed workload per domain (arch-specific, see
`oracle_test.go`), and every solver test carries a residual or an
exact-reference check. A benchmark result without the gates green is
not a result.
## Reports
One live report exists: `docs/benchmarks/release-vs-head.md`, the
newest release tag against the working tree. Regenerate it with:
```sh
just bench-report
```
The recipe checks out the latest `v*` tag in a scratch worktree, runs
the representative set `bench_set` names over four interleaved rounds
on both revisions, and rewrites the file. The set holds one benchmark
per kernel family; a name either revision lacks is left out rather
than counted. Run it on an idle machine and commit the file it writes
alongside the release it describes. Per-change exploration numbers
belong in the commit's own review, not in a growing pile of report
files; the repository carries the one comparison that matters, the
release a reader has against the tree as it stands.
+204
View File
@@ -0,0 +1,204 @@
# Development
How to work on Tensor.
## Prerequisites
- Go 1.27.1, the exact version `go.mod` declares and the newest
stable release at the time of writing. Verify the installed version
against the release list rather than memory: `go version`.
- [just](https://github.com/casey/just) for the recipes. Three of
them (`test`, `fmt-check`, `fuzz-all`) are Perl scripts.
- Perl, for those recipes and for the CI steps that carry logic. Only
the interpreter's own builtins are used, so no module installation
is needed.
- A C compiler (`gcc`) for the race detector, which `just race` and
`just gates` run; `-race` requires cgo.
Nothing else: Tensor has zero third-party dependencies.
## Setup
```sh
git clone https://sourcedock.dev/petrbalvin/tensor.git
cd tensor
just build
just test
```
## Recipes
Every recipe in the project's file, and what it does. Taken from the
file itself, so the names and the list match it exactly.
| Recipe | What it does |
|---|---|
| `just` | lists the recipes |
| `just build` | compiles everything, the example programs included |
| `just test` | the test gate: the full suite with no cache, the coverage profile and the 80 percent floor |
| `just race` | the same suite under the race detector; the expensive one, run once per task by `just gates` |
| `just unit ./internal/core/ TestName` | a fast scoped run for iterating: cached, no race, no coverage |
| `just fuzz FuzzName ./io 60s` | a time-boxed fuzz of one target in one package; never a gate |
| `just bench` | benchmarks with `-benchmem`, five counts; on an idle machine only |
| `just fmt` | formats all Go sources in place with `gofmt` |
| `just fmt-check` | verifies that `gofmt` produces no diff; prints nothing on success |
| `just vet` | `go vet` and `go fix -diff` |
| `just gates` | the definition of done in one command: `build`, `fmt-check`, `vet`, `test`, `race`, in that order |
| `just clean` | removes the build artefacts (`bin/`, `coverage.out`) |
| `just docs-check` | runs every Go program in `README.md` from a temporary module, so the documentation cannot claim what the code no longer does |
| `just fuzz-all 5s` | fuzzes every target for the budget each; exploration, never a gate |
`docs-check` and `fuzz-all` are the project extensions; none of
them is a gate.
The `packages` value behind `test`, `race`, `unit` and `bench` names
the logic packages and leaves `examples/` out: those are main
programs with no tests, and the build is what compiles them. Tensor
is a library, so the binary recipes (`install`, `run`, `dev`) have no
referent here and are absent from the file.
The scripted recipes keep their logic in Perl rather than in the
shell, which is the repository rule for every non-product script: the
shell starts commands, and anything with a branch or a loop is Perl
using the interpreter's own builtins.
`gofmt` is the single formatting authority: there is no configuration
beyond it, `just fmt-check` is the gate and `just fmt` the fix.
## Running a single test
```sh
just unit ./internal/core/ TestQuo
```
`unit` is the scoped, cached run for iterating; the second argument
is a regular expression matched against test names. Combine with
`-v` for the sub-test names, or call `go test` directly:
```sh
go test -run TestQuo -v -count=1 ./internal/core/
```
`-count=1` defeats the test cache when a result looks stale.
The runnable documentation is part of the suite, so it is exercised
the same way. Each package carries its examples beside its tests:
```sh
go test ./linalg/ -run Example -count=1 -v
```
A godoc example that stops compiling, or whose printed output drifts
from its `// Output:` comment, fails the suite rather than the
reader. The programs in `README.md` are checked the same way, though
outside the suite, because they are whole `main` programs:
```sh
just docs-check
```
which extracts every `go` block into a temporary module against the
working tree, runs it, and reports the block that failed.
## Coverage
```sh
just test
go tool cover -func=coverage.out
```
The `total:` line is the number that matters, and `just test` fails
below 80 percent. The sweep names the logic packages, so every
library package is measured while the examples stay out of the
denominator. For an HTML report:
```sh
go tool cover -html=coverage.out -o coverage.html
```
Two harnesses inside the suite guard properties that coverage
percentages do not describe, and both live at the root:
- **`TestOracle`** pins a raw-bit digest of one fixed workload per
domain, per platform and per build. A digest that moves is either a
deliberate arithmetic change or a regression, and the difference is
decided by the person who moved it, not by the test.
- **`TestNoResourceLeaks`** measures the heap across three blocks of
ten rounds and fails on a net rise above 256 KiB, which is how a
buffer that stops being released is caught before it becomes an
outage.
## Benchmarks
```sh
just bench
```
One package at a time, with a fixed budget:
```sh
go test ./internal/core/ -bench 'BenchmarkMatMul$' -benchtime 2s -run xxx
```
Benchmark on an idle machine, compare only runs made in one process
against each other, and treat a few percent as noise. The packages
carry 147 benchmarks, and the weight sits where the time is: 84 in
`internal/core`, 20 in `signal`, 13 in `stats`, 11 in `integrate`, 8
in `optim`, 7 in `linalg`, 3 in `grad` and 1 in `internal/engine`.
The binding measurement method, the report template and the measured
reports live in [docs/BENCHMARKING.md](BENCHMARKING.md) and
[docs/benchmarks/](benchmarks/).
## Debugging the build
```sh
go build -gcflags='-m' ./internal/core/ # inlining decisions
go build -gcflags='-S' ./internal/core/ # what the compiler generated
```
There is exactly one build, and it is the product:
| Build | Command | Assumes |
|---|---|---|
| portable | `go build ./...` | the toolchain default code generation, no pinned `GOAMD64` |
The portable build pins no `GOAMD64` level: the compiler has no
auto-vectoriser, so a pinned higher level would buy only scalar FMA
contraction, which the bit-pinned kernels suppress by spelling anyway
(`float64(a*b) + c`). A build pinned to a level the oracle has no
digest block for skips loudly, so a quiet mismatch cannot happen.
## Continuous integration
Gitea Actions workflows live in `.gitea/workflows/`, are written by
hand, and enforce the same gate set as `just gates`, with scripted
steps in Perl and parallelism bounded to the shared runner box:
- **`test.yml`**, on every push and pull request to `development`:
build, format check, vet, the full suite with the coverage floor,
and the oracle digests for the platform. Race is absent on purpose:
the shared box cannot afford it on every push. The one-iteration
benchmark smoke that once rode along is retired outright: the
minimum degree battery's 3-D mesh scan alone runs for minutes on one
core and allocates terabytes cumulatively, so no form of it fits the
shared box, and benchmarking is deliberate work on a developer
machine.
- **`race.yml`**, dispatched by hand: the suite under the race
detector, with the oracle digests across `fedora`, `alpine` and
`openeuler`, which is the glibc against musl check the
floating-point kernels need.
- **`release.yml`**, on a `v*` tag: the gate set minus race once at
the tag, then the Gitea release created from the matching
`CHANGELOG.md` section.
A green `just gates` locally is the fastest way to a green pipeline.
## Releases
Releases are cut by merging `development` into `main` and tagging
`vX.Y.Z`. The tag pipeline runs the gates at the tag and publishes
the release with the CHANGELOG section as its notes: the pipeline
reads the section that begins at `## [X.Y.Z]` and stops at the next
`## [`, and refuses a tag whose section is missing or empty. Nothing
is injected into the build; the toolchain records the tag because the
build simply happens there. Before cutting a tag, run `just gates`
locally: the local gate is the one that races the tree.
File diff suppressed because it is too large Load Diff
+308
View File
@@ -0,0 +1,308 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package tensor_test
// The godoc examples: every flagship workflow as a runnable, checked
// snippet. pkg.go.dev renders these beside the API, and `go test`
// executes them, so the documentation cannot rot.
import (
"fmt"
"log"
"math"
tensor "sourcedock.dev/petrbalvin/tensor"
grad "sourcedock.dev/petrbalvin/tensor/grad"
)
// A rank-1 array from literals, element-wise arithmetic, a reduction.
func ExampleAdd() {
a, err := tensor.FromFloats([]float64{1, 2, 3, 4}, 4)
if err != nil {
log.Fatal(err)
}
b, _ := tensor.FromFloats([]float64{10, 20, 30, 40}, 4)
sum, _ := tensor.Add(a, b)
mean, _ := tensor.Mean(sum)
fmt.Println(sum, mean)
// Output: float (4) [11, 22, 33, 44] 27.5
}
// The 2-D matrix product.
func ExampleMatMul2D() {
a, _ := tensor.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
b, _ := tensor.FromFloats([]float64{5, 6, 7, 8}, 2, 2)
prod, err := tensor.MatMul2D(a, b)
if err != nil {
log.Fatal(err)
}
fmt.Println(prod)
// Output: float (2, 2) [19, 22, 43, 50]
}
// Einstein summation, batched over an ellipsis axis.
func ExampleEinsum() {
a, _ := tensor.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2)
b, _ := tensor.FromFloats([]float64{1, 0, 0, 1, 1, 0, 0, 1}, 2, 2, 2)
// Batched matrix product: each batch times its identity.
got, err := tensor.Einsum("...ij,...jk->...ik", a, b)
if err != nil {
log.Fatal(err)
}
fmt.Println(got.Shape(), got.FloatAt(0), got.FloatAt(3))
// Output: [2 2 2] 1 4
}
// Reverse-mode autograd: the gradient of Σ (x·w)² arrives exact.
func ExampleTensor_Backward() {
x, _ := grad.FromFloat64s([]float64{1, 2, 3}, true, 3)
w, _ := grad.FromFloat64s([]float64{0.5, -1, 2}, false, 3)
prod, _ := x.Mul(w)
sq, _ := prod.Pow(2)
loss, _ := sq.Sum()
if err := loss.Backward(); err != nil {
log.Fatal(err)
}
// dL/dx = 2·x·w², non-negative because w enters squared.
g := x.Grad()
fmt.Printf("%.4g %.4g %.4g\n", g.FloatAt(0), g.FloatAt(1), g.FloatAt(2))
// Output: 0.5 4 24
}
// Ordinary least squares with the full classical inference.
func ExampleLinearRegression() {
// y = 2 + 3x on x = 0..5, the intercept column first.
design, _ := tensor.FromFloats([]float64{
1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5,
}, 6, 2)
y, _ := tensor.FromFloats([]float64{2.1, 4.9, 8.2, 11.1, 13.8, 17.2}, 6)
res, err := tensor.LinearRegression(design, y)
if err != nil {
log.Fatal(err)
}
fmt.Printf("slope %.3f ± %.3f, R2 %.4f\n",
res.Coefficients[1], res.StandardErrors[1], res.RSquared)
// Output: slope 3.003 ± 0.044, R2 0.9991
}
// Solving an initial value problem: the exponential decay y' = −y.
func ExampleIntegrateODE() {
f := func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.MulF(y, -1), nil
}
y0, _ := tensor.FromFloats([]float64{1}, 1)
end, err := tensor.IntegrateODE(f, 0, 1, y0, tensor.ODEOptions{})
if err != nil {
log.Fatal(err)
}
fmt.Printf("y(1) = %.6f\n", end.FloatAt(0))
// Output: y(1) = 0.367880
}
// The globally adaptive cubature over a box: the 2-D Gaussian.
func ExampleIntegrateND() {
got, err := tensor.IntegrateND(func(x []float64) float64 {
return math.Exp(-x[0]*x[0] - x[1]*x[1])
}, []float64{-3, -3}, []float64{3, 3}, tensor.CubatureOptions{Tolerance: 1e-11})
if err != nil {
log.Fatal(err)
}
fmt.Printf("%.6f\n", got)
// Output: 3.141454
}
// The Crank-Nicolson heat equation holds its eigenmode shape while it
// decays.
func ExampleIntegrateHeat1D() {
const (
n = 49
kappa = 0.1
dx = 1.0 / 50
)
u0 := make([]float64, n)
for i := range n {
u0[i] = math.Sin(math.Pi * float64(i+1) * dx)
}
u0Arr, _ := tensor.FromFloats(u0, n)
states, err := tensor.IntegrateHeat1D(u0Arr, kappa, dx, 1.0, 0.001, 2, 0, 0)
if err != nil {
log.Fatal(err)
}
last := (states.Shape()[0] - 1) * n
// The mode's centre started at 1 and decays by exp(−κπ²t).
decay := math.Exp(-kappa * math.Pi * math.Pi)
fmt.Printf("decay %.4f, centre %.4f\n", decay, states.FloatAt(last+n/2))
// Output: decay 0.3727, centre 0.3728
}
// The Lanczos eigensolver on a sparse symmetric matrix.
func ExampleSpEigen() {
dense, _ := tensor.FromFloats([]float64{2, 1, 1, 2}, 2, 2)
sp, _ := tensor.SparseFrom(dense)
vals, _, err := tensor.SpEigen(sp, 2, tensor.NewGenerator(7))
if err != nil {
log.Fatal(err)
}
fmt.Printf("%.6f %.6f\n", vals.FloatAt(0), vals.FloatAt(1))
// Output: 3.000000 1.000000
}
// The complete elliptic integral of the first kind by the AGM.
func ExampleEllipticK() {
m, _ := tensor.FromFloats([]float64{0, 0.5}, 2)
k, err := tensor.EllipticK(m)
if err != nil {
log.Fatal(err)
}
fmt.Printf("%.10f %.10f\n", k.FloatAt(0), k.FloatAt(1))
// Output: 1.5707963268 1.8540746773
}
// Haar wavelets turn an off-grid step into one detail coefficient,
// and the inverse restores the signal exactly.
func ExampleDWT() {
vals := make([]float64, 16)
for i := 5; i < 16; i++ {
vals[i] = 1
}
x, _ := tensor.FromFloats(vals, 16)
c, err := tensor.DWT(x, 1)
if err != nil {
log.Fatal(err)
}
restored, _ := tensor.IDWT(c, 1)
ok := math.Abs(restored.FloatAt(7)-1) < 1e-12 && restored.FloatAt(4) == 0
fmt.Printf("detail[2] = %.4f, restored: %v\n", c.FloatAt(10), ok)
// Output: detail[2] = -0.7071, restored: true
}
// The autocorrelation of a periodic signal is itself periodic.
func ExampleAutocorrelate() {
vals := make([]float64, 64)
for i := range vals {
vals[i] = math.Cos(2 * math.Pi * float64(i) / 16)
}
x, _ := tensor.FromFloats(vals, 64)
acf, err := tensor.Autocorrelate(x, 16)
if err != nil {
log.Fatal(err)
}
// The biased normalisation tapers the tail: 48/64 of the pairs
// overlap at lag 16.
fmt.Printf("acf(0) = %.4f, acf(16) = %.4f\n", acf.FloatAt(0), acf.FloatAt(16))
// Output: acf(0) = 1.0000, acf(16) = 0.7500
}
// Differential evolution crosses Rastrigin's minefield of local
// minima to the global one.
func ExampleMinimiseDifferentialEvolution() {
lower, _ := tensor.FromFloats([]float64{-5.12, -5.12}, 2)
upper, _ := tensor.FromFloats([]float64{5.12, 5.12}, 2)
_, fv, err := tensor.MinimiseDifferentialEvolution(func(a *tensor.Array) (float64, error) {
s := 0.0
for i := range 2 {
z := a.FloatAt(i)
s += z*z - 10*math.Cos(2*math.Pi*z) + 10
}
return s, nil
}, lower, upper, tensor.DifferentialEvolutionOptions{Generations: 800})
if err != nil {
log.Fatal(err)
}
fmt.Printf("%.1e\n", fv)
// Output: 0.0e+00
}
// Slicing a contiguous selection returns a read-only view sharing the
// storage; an interior range is copied.
func ExampleSlice() {
m, _ := tensor.FromFloats([]float64{
1, 2, 3, 4,
5, 6, 7, 8,
9, 10, 11, 12,
}, 3, 4)
// Whole rows: a view.
rows, err := tensor.Slice(m, 0, 1, 3)
if err != nil {
log.Fatal(err)
}
// An interior column range: a copy.
block, _ := tensor.Slice(rows, 1, 1, 3)
fmt.Println(rows.Shape(), rows.FloatAt(0), block.Shape(), block.FloatAt(0))
// Output: [2 4] 5 [2 2] 6
}
// A size-1 dimension replicates over the target, and anything that
// would not is an error rather than a silent broadcast.
func ExampleBroadcastTo() {
col, _ := tensor.FromFloats([]float64{1, 2, 3}, 3, 1)
wide, err := tensor.BroadcastTo(col, 3, 4)
if err != nil {
log.Fatal(err)
}
_, bad := tensor.BroadcastTo(col, 2, 4)
fmt.Println(wide)
fmt.Println("refused:", bad != nil)
// Output: float (3, 4) [1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3]
// refused: true
}
// The gamma function, including the negative half line.
func ExampleGamma() {
x, _ := tensor.FromFloats([]float64{-0.5, 0.5, 5}, 3)
g, err := tensor.Gamma(x)
if err != nil {
log.Fatal(err)
}
fmt.Printf("%.7f %.7f %.1f\n", g.FloatAt(0), g.FloatAt(1), g.FloatAt(2))
// Output: -3.5449077 1.7724539 24.0
}
// Linear interpolation clamps outside the sampled range; the monotone
// cubic passes through the same knots with a bounded slope.
func ExampleInterpolate() {
xs, _ := tensor.FromFloats([]float64{0, 1, 3}, 3)
ys, _ := tensor.FromFloats([]float64{0, 2, 2.5}, 3)
query, _ := tensor.FromFloats([]float64{0.5, 2, 5}, 3)
lin, err := tensor.Interpolate(xs, ys, query)
if err != nil {
log.Fatal(err)
}
mono, _ := tensor.InterpolateMonotone(xs, ys, query)
// Both stay inside the bracketing samples, and the tail clamps to
// the last knot.
fmt.Printf("linear %.4f %.4f %.4f\n", lin.FloatAt(0), lin.FloatAt(1), lin.FloatAt(2))
fmt.Printf("pchip %.4f %.4f %.4f\n", mono.FloatAt(0), mono.FloatAt(1), mono.FloatAt(2))
// Output: linear 1.0000 2.2500 2.5000
// pchip 1.2621 2.3716 2.5000
}
// The Halton sequence stratifies progressively in every base; the
// index-zero origin is dropped, as it carries no information.
func ExampleHaltonPoints() {
pts, err := tensor.HaltonPoints(3, 2, 0)
if err != nil {
log.Fatal(err)
}
fmt.Println(pts.Shape())
for i := range 3 {
fmt.Printf("(%.4f, %.4f) ", pts.FloatAt(2*i), pts.FloatAt(2*i+1))
}
fmt.Println()
// Output: [3 2]
// (0.5000, 0.3333) (0.2500, 0.6667) (0.7500, 0.1111)
}
// A fixed seed replays the same stream, which is what makes a random
// workflow testable.
func ExampleNewGenerator() {
u, err := tensor.Floats(tensor.NewGenerator(7), 3)
if err != nil {
log.Fatal(err)
}
again, _ := tensor.Floats(tensor.NewGenerator(7), 3)
fmt.Printf("%.6f %.6f %.6f, equal: %v\n",
u.FloatAt(0), u.FloatAt(1), u.FloatAt(2), tensor.Equal(u, again))
// Output: 0.055360 0.172116 0.717576, equal: true
}
+165
View File
@@ -0,0 +1,165 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command deconv recovers a sharp image from a blurred, noisy
// observation by gradient descent through the Fourier transform: the
// convolution runs as a spectral product, the loss differentiates
// through FFT2 and IFFT2, and Tikhonov regularisation keeps the noise
// from winning. Astronomical PSF deconvolution in a page of gradient,
// the workload the spectral autograd exists for.
//
// Usage: go run ./examples/deconv
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
n = 32
sigma = 6.0
)
// The truth: one off-centre Gaussian source.
truth := make([]float64, n*n)
for r := range n {
for c := range n {
dx := float64(c) - 20.0
dy := float64(r) - 12.0
truth[r*n+c] = math.Exp(-(dx*dx + dy*dy) / (2 * 2.5 * 2.5))
}
}
// The PSF: a wider Gaussian, the blur to undo.
psf := make([]float64, n*n)
for r := range n {
for c := range n {
dx := float64(c) - n/2
dy := float64(r) - n/2
psf[r*n+c] = math.Exp(-(dx*dx + dy*dy) / (2 * sigma * sigma))
}
}
obsArr, err := tensor.FromFloats(truth, n, n)
if err != nil {
log.Fatal(err)
}
psfArr, err := tensor.FromFloats(psf, n, n)
if err != nil {
log.Fatal(err)
}
// The blur runs as a spectral product; the constants ride the same
// graph nodes without requiring grad.
kT := tensor.FromArray(psfArr, false)
kF, err := kT.FFT2()
if err != nil {
log.Fatal(err)
}
obsT := tensor.FromArray(obsArr, false)
spec, err := obsT.FFT2()
if err != nil {
log.Fatal(err)
}
prod, err := spec.Mul(kF)
if err != nil {
log.Fatal(err)
}
blurred, err := prod.IFFT2()
if err != nil {
log.Fatal(err)
}
noise := tensor.NewGenerator(42)
noisy := make([]float64, n*n)
for i := range noisy {
noisy[i] = real(blurred.Data().ComplexAt(i)) + 0.01*noise.NormalUnit()
}
obsC, err := tensor.FromFloats(noisy, n, n)
if err != nil {
log.Fatal(err)
}
obsComplex, err := tensor.Astype(obsC, tensor.Complex)
if err != nil {
log.Fatal(err)
}
obsTt := tensor.FromArray(obsComplex, false)
// The spectral loss is a stiff quadratic (curvature ~ |K|² per
// frequency), exactly the landscape plain gradient descent crawls
// on and Newton-CG eats: the CG solve rides the Hessian-vector
// product through the same FFT chain, two backward passes per
// iteration, no dense Hessian ever formed.
xArr, err := tensor.FromFloats(make([]float64, n*n), n, n)
if err != nil {
log.Fatal(err)
}
objective := func(x *tensor.Tensor) (*tensor.Tensor, error) {
spec, err := x.FFT2()
if err != nil {
return nil, err
}
prod, err := spec.Mul(kF)
if err != nil {
return nil, err
}
model, err := prod.IFFT2()
if err != nil {
return nil, err
}
resid, err := model.Sub(obsTt)
if err != nil {
return nil, err
}
dataTerm, err := resid.Abs2()
if err != nil {
return nil, err
}
regTerm, err := x.Abs2()
if err != nil {
return nil, err
}
reg, err := regTerm.Scale(1e-3)
if err != nil {
return nil, err
}
both, err := dataTerm.Add(reg)
if err != nil {
return nil, err
}
return both.Sum()
}
solution, loss, err := tensor.MinimiseNewtonCG(objective, xArr,
tensor.NewtonCGOptions{Tolerance: 1e-8, MaxIterations: 80})
if err != nil {
log.Fatal(err)
}
xArr = solution
fmt.Printf("newton-cg converged, loss = %.6e\n", loss)
peakOf := func(a *tensor.Array) (int, int, float64) {
best := math.Inf(-1)
br, bc := 0, 0
for r := range n {
for c := range n {
if v := a.FloatAt(r*n + c); v > best {
best, br, bc = v, r, c
}
}
}
return br, bc, best
}
tr, tc, tv := peakOf(obsArr)
br, bc, bv := peakOf(obsC)
rr, rc, rv := peakOf(xArr)
fmt.Printf("truth peak at (%d, %d), height %.3f\n", tr, tc, tv)
fmt.Printf("blurred observation peak at (%d, %d), height %.3f\n", br, bc, bv)
fmt.Printf("recovered peak at (%d, %d), height %.3f (final loss %.3e)\n", rr, rc, rv, loss)
if rr != tr || rc != tc {
log.Fatal("the recovery did not localise the source")
}
if math.Abs(rv-tv) > 0.25*tv {
log.Fatal("the recovery did not restore the source height")
}
}
+62
View File
@@ -0,0 +1,62 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command fft demonstrates the Fourier transform: it synthesises a
// signal from two sinusoids, transforms it, and prints the dominant
// frequency components.
//
// Usage: go run ./examples/fft
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
sig "sourcedock.dev/petrbalvin/tensor/signal"
)
func main() {
// Two sinusoids: 5 Hz and 13 Hz, sampled at 100 Hz for 2 seconds.
const (
fs = 100.0
seconds = 2.0
)
n := int(fs * seconds)
vals := make([]float64, n)
for i := range n {
t := float64(i) / fs
vals[i] = math.Sin(2*math.Pi*5*t) + 0.5*math.Sin(2*math.Pi*13*t)
}
signal, err := tensor.FromFloats(vals, n)
if err != nil {
log.Fatal(err)
}
spec, err := sig.FFT(signal)
if err != nil {
log.Fatal(err)
}
freqs := sig.FFTFreq(n, 1/fs)
// Find the two strongest bins (excluding DC).
var peaks [2]struct {
freq float64
mag float64
}
for i := 1; i < n/2; i++ {
m, _ := tensor.ComplexAt(spec, i)
mag := math.Hypot(real(m), imag(m))
for p := range peaks {
if mag > peaks[p].mag {
peaks[p].freq, _ = tensor.FloatAt(freqs, i)
peaks[p].mag = mag
break
}
}
}
for _, p := range peaks {
fmt.Printf("peak at %.1f Hz (magnitude %.1f)\n", p.freq, p.mag)
}
}
+98
View File
@@ -0,0 +1,98 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command fits builds a synthetic star field, saves it as a FITS
// primary image, reads it back and recovers the brightest star's
// position by a centre-of-mass centroid, the first step of any
// aperture photometry pipeline.
//
// Usage: go run ./examples/fits
package main
import (
"fmt"
"log"
"math"
"os"
"path/filepath"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
size = 128
sigma = 2.0 // pixels, the seeing disk
)
// Three stars of different brightness on a flat sky background.
type star struct {
x, y, flux float64
}
stars := []star{
{40.5, 60.5, 900},
{80.5, 30.5, 300},
{95.5, 95.5, 120},
}
field := make([]float64, size*size)
for i := range size {
for j := range size {
v := 100.0 // sky
for _, s := range stars {
d2 := (float64(i)-s.y)*(float64(i)-s.y) + (float64(j)-s.x)*(float64(j)-s.x)
v += s.flux * math.Exp(-d2/(2*sigma*sigma))
}
field[i*size+j] = v
}
}
img, err := tensor.FromFloats(field, size, size)
if err != nil {
log.Fatal(err)
}
path := filepath.Join(os.TempDir(), "tensor-example-stars.fits")
defer os.Remove(path)
headers := map[string]string{
"OBJECT": "synthetic field",
"EXPTIME": "30",
"FILTER": "V",
}
if err := tensor.SaveFITS(path, img, headers); err != nil {
log.Fatal(err)
}
back, hdr, err := tensor.LoadFITS(path)
if err != nil {
log.Fatal(err)
}
fmt.Printf("wrote and read %s\n", filepath.Base(path))
for _, k := range []string{"OBJECT", "EXPTIME", "FILTER"} {
fmt.Printf(" %s = %s\n", k, hdr[k])
}
if back.Shape()[0] != size || back.Shape()[1] != size {
log.Fatalf("round trip changed the shape: %v", back.Shape())
}
// Locate the brightest pixel, then centroid a 9x9 window around
// it with the sky level subtracted.
best, bestVal := 0, -1.0
for i := range size * size {
if v := back.FloatAt(i); v > bestVal {
best, bestVal = i, v
}
}
by, bx := best/size, best%size
sum, sx, sy := 0.0, 0.0, 0.0
for i := by - 4; i <= by+4; i++ {
for j := bx - 4; j <= bx+4; j++ {
w := back.FloatAt(i*size+j) - 100
if w < 0 {
w = 0
}
sum += w
sx += w * float64(j)
sy += w * float64(i)
}
}
fmt.Printf("\nbrightest star: peak at (x=%d, y=%d), %.0f counts\n", bx, by, bestVal)
fmt.Printf("centroid of the 9x9 window: (x=%.2f, y=%.2f)\n", sx/sum, sy/sum)
fmt.Println("true position: (x=40.50, y=60.50)")
}
+242
View File
@@ -0,0 +1,242 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command helmholtz solves the discretised Helmholtz equation, the
// backbone of frequency-domain electromagnetics, in both of its
// solver shapes. The time-harmonic wave equation
//
// -∇²ψ - k²ψ = f
//
// on a 2-D grid gives a complex symmetric (non-Hermitian) sparse
// system, which the BiCGSTAB solver handles. Adding a small imaginary
// part to k², the way a lossy medium does, makes the operator
// Hermitian positive-definite and the conjugate gradient solver
// applies. Both solutions are verified against the dense solve, and
// the Hermitian operator's resonant modes come from the sparse
// eigensolver.
//
// Usage: go run ./examples/helmholtz
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
const grid = 24 // interior points per side
// laplacianCOO assembles the 5-point discrete -∇² on the interior of
// a grid*grid domain with Dirichlet walls, one entry per stencil
// point. The value at (i,j) is k2 times the identity there.
func helmholtzCOO(k2 complex128) (*tensor.SparseCOO, int) {
n := grid * grid
var idx []int64
var val []complex128
at := func(i, j int) int { return i*grid + j }
for i := range grid {
for j := range grid {
p := at(i, j)
// 4/h² on the diagonal with h = 1 in grid units, minus k².
idx = append(idx, int64(p), int64(p))
val = append(val, 4-k2)
if i > 0 {
idx = append(idx, int64(p), int64(at(i-1, j)))
val = append(val, -1)
}
if i < grid-1 {
idx = append(idx, int64(p), int64(at(i+1, j)))
val = append(val, -1)
}
if j > 0 {
idx = append(idx, int64(p), int64(at(i, j-1)))
val = append(val, -1)
}
if j < grid-1 {
idx = append(idx, int64(p), int64(at(i, j+1)))
val = append(val, -1)
}
}
}
indices, err := tensor.FromInts(idx, len(val), 2)
if err != nil {
log.Fatal(err)
}
values, err := tensor.FromComplexes(val, len(val))
if err != nil {
log.Fatal(err)
}
coo, err := tensor.NewSparseCOO(indices, values, []int{n, n})
if err != nil {
log.Fatal(err)
}
return coo, n
}
// landauCOO assembles the Hamiltonian of a charged particle on the
// same grid threading a perpendicular magnetic field, the Peierls
// substitution: every hop carries the phase the vector potential
// gives it, forward and conjugate backward, so the operator stays
// Hermitian. A positive mass term m² makes it positive-definite.
func landauCOO(m2, flux float64) *tensor.SparseCOO {
n := grid * grid
var idx []int64
var val []complex128
at := func(i, j int) int { return i*grid + j }
phase := func(i int) float64 { return 2 * math.Pi * flux * float64(i) }
for i := range grid {
for j := range grid {
p := at(i, j)
idx = append(idx, int64(p), int64(p))
val = append(val, complex(4+m2, 0))
if i > 0 {
idx = append(idx, int64(p), int64(at(i-1, j)))
val = append(val, -1+0i)
}
if i < grid-1 {
idx = append(idx, int64(p), int64(at(i+1, j)))
val = append(val, -1+0i)
}
if j > 0 {
idx = append(idx, int64(p), int64(at(i, j-1)))
val = append(val, -complex(math.Cos(phase(i)), math.Sin(phase(i))))
}
if j < grid-1 {
idx = append(idx, int64(p), int64(at(i, j+1)))
val = append(val, -complex(math.Cos(phase(i)), -math.Sin(phase(i))))
}
}
}
indices, err := tensor.FromInts(idx, len(val), 2)
if err != nil {
log.Fatal(err)
}
values, err := tensor.FromComplexes(val, len(val))
if err != nil {
log.Fatal(err)
}
coo, err := tensor.NewSparseCOO(indices, values, []int{n, n})
if err != nil {
log.Fatal(err)
}
return coo
}
// source is a point drive at the grid centre, the field of a small
// antenna.
func source(n int) *tensor.Array {
rhs := make([]complex128, n)
rhs[(grid/2)*grid+grid/2] = 1 + 0i
b, err := tensor.FromComplexes(rhs, n)
if err != nil {
log.Fatal(err)
}
return b
}
// residual returns ||b - A·x||₂ by reassembling A densely, the ground
// truth the sparse solver is checked against.
func residual(a *tensor.SparseCOO, x, b *tensor.Array, n int) float64 {
dense, err := a.Dense()
if err != nil {
log.Fatal(err)
}
ax, err := tensor.MatMul2D(dense, x)
if err != nil {
log.Fatal(err)
}
worst := 0.0
for i := range n {
av, err := tensor.ComplexAt(ax, i)
if err != nil {
log.Fatal(err)
}
bv, err := tensor.ComplexAt(b, i)
if err != nil {
log.Fatal(err)
}
if d := math.Hypot(real(av-bv), imag(av-bv)); d > worst {
worst = d
}
}
return worst
}
func main() {
const n = grid * grid
// A propagating mode: k = 2.5 in grid units, safely away from the
// discrete resonances at k² = 2-2cos(p*pi/(grid+1)).
k := 2.5 + 0i
a, _ := helmholtzCOO(k * k)
b := source(n)
x, err := tensor.SpSolveComplexBiCGSTAB(a, b, 1e-12, 2000)
if err != nil {
log.Fatal(err)
}
fmt.Println("lossless Helmholtz system, -nabla^2 - k^2, k = 2.5")
fmt.Printf(" unknowns: %d, stored nonzeros: %d\n", n, len(a.Values.RawComplexes()))
fmt.Printf(" BiCGSTAB residual ||b - A x|| = %.3g\n", residual(a, x, b, n))
// The genuinely Hermitian complex problem: a charged particle on
// the same grid in a perpendicular magnetic field. The Peierls
// phases make every hop complex, the forward and backward hop
// conjugates of each other, so the operator is Hermitian, and the
// mass term keeps it positive-definite: exactly the shape the
// conjugate gradient solver wants.
h := landauCOO(1.0, 1.0/25)
xh, err := tensor.SpSolveComplexCG(h, b, 1e-12, 2000)
if err != nil {
log.Fatal(err)
}
fmt.Println("\nLandau Hamiltonian on the grid, mass^2 = 1, flux 1/25 (Hermitian positive-definite)")
fmt.Printf(" CG residual ||b - A x|| = %.3g\n", residual(h, xh, b, n))
// Resonant modes of the lossless cavity: the largest eigenvalues
// of the discrete negative Laplacian are the highest-Q modes.
lap, _ := helmholtzCOO(0)
vals, vecs, err := tensor.SpEigenComplex(lap, 3, tensor.NewGenerator(4))
if err != nil {
log.Fatal(err)
}
fmt.Println("\ncavity modes: largest eigenvalues of -nabla^2")
for j := range 3 {
lam, err := tensor.FloatAt(vals, j)
if err != nil {
log.Fatal(err)
}
// Verify each Ritz pair: ||A v - lambda v|| must be small.
vcol, err := tensor.Slice(vecs, 1, j, j+1)
if err != nil {
log.Fatal(err)
}
av, err := tensor.MatMul2D(mustDense(lap), vcol)
if err != nil {
log.Fatal(err)
}
worst := 0.0
for i := range n {
a1, _ := tensor.ComplexAt(av, i)
v1, _ := tensor.ComplexAt(vcol, i)
if d := math.Hypot(real(a1-complex(lam, 0)*v1), imag(a1-complex(lam, 0)*v1)); d > worst {
worst = d
}
}
fmt.Printf(" lambda = %8.4f, residual %.3g\n", lam, worst)
}
// The analytic eigenvalues of the grid Laplacian are
// 4-2cos(p*pi/(grid+1))-2cos(q*pi/(grid+1)); the largest is p = q =
// grid, where both cosines approach -1 and the value nears 8.
fmt.Println(" (analytic maximum: 4 - 4cos(24pi/25) = 7.9685)")
}
func mustDense(a *tensor.SparseCOO) *tensor.Array {
d, err := a.Dense()
if err != nil {
log.Fatal(err)
}
return d
}
+102
View File
@@ -0,0 +1,102 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command hmc samples a correlated two-dimensional Gaussian by
// Hamiltonian Monte Carlo on the differentiable log density, and
// checks the chain against the distribution's known moments: mean
// zero, unit variances, correlation 0.9. The sampler never sees an
// analytic gradient, only the autograd's.
//
// Usage: go run ./examples/hmc
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const rho = 0.9
// Precision matrix of the correlated Gaussian (up to scale, which
// the unnormalised density does not need).
prec, err := tensor.FromFloats([]float64{1, -rho, -rho, 1}, 2, 2)
if err != nil {
log.Fatal(err)
}
aT := tensor.FromArray(prec, false)
logDensity := func(q *tensor.Tensor) (*tensor.Tensor, error) {
// log p(q) ∝ −½ qᵀAq: one matrix-vector product on the graph,
// then the inner product with q itself.
r, err := aT.MatMul(q)
if err != nil {
return nil, err
}
quad, err := r.Mul(q)
if err != nil {
return nil, err
}
s, err := quad.Sum()
if err != nil {
return nil, err
}
return s.Scale(-0.5)
}
q0, err := tensor.FromFloats([]float64{0.5, -0.5}, 2)
if err != nil {
log.Fatal(err)
}
samples, err := tensor.SampleHMC(logDensity, q0, tensor.HMCOptions{
Step: 0.15,
Steps: 12,
BurnIn: 500,
Thin: 5,
Samples: 20000,
Seed: 7,
})
if err != nil {
log.Fatal(err)
}
n := samples.Shape()[0]
mean0, mean1 := 0.0, 0.0
for i := range n {
mean0 += samples.FloatAt(i * 2)
mean1 += samples.FloatAt(i*2 + 1)
}
mean0 /= float64(n)
mean1 /= float64(n)
var0, var1, cov := 0.0, 0.0, 0.0
for i := range n {
d0 := samples.FloatAt(i*2) - mean0
d1 := samples.FloatAt(i*2+1) - mean1
var0 += d0 * d0
var1 += d1 * d1
cov += d0 * d1
}
var0 /= float64(n - 1)
var1 /= float64(n - 1)
cov /= float64(n - 1)
corr := cov / math.Sqrt(var0*var1)
fmt.Printf("samples %d\n", n)
fmt.Printf("mean (%.3f, %.3f), want (0, 0)\n", mean0, mean1)
// The covariance is A⁻¹: unit-over-(1−ρ²) variances around the
// correlation rho.
wantVar := 1 / (1 - rho*rho)
fmt.Printf("variance (%.3f, %.3f), want (%.3f, %.3f)\n", var0, var1, wantVar, wantVar)
fmt.Printf("correlation %.3f, want %.3f\n", corr, rho)
if math.Abs(mean0) > 0.05 || math.Abs(mean1) > 0.05 {
log.Fatal("the sample mean drifted")
}
if math.Abs(corr-rho) > 0.02 {
log.Fatalf("the sample correlation %.3f missed %.3f", corr, rho)
}
if math.Abs(var0-wantVar) > 0.15*wantVar || math.Abs(var1-wantVar) > 0.15*wantVar {
log.Fatalf("the sample variances (%.3f, %.3f) missed %.3f", var0, var1, wantVar)
}
}
+108
View File
@@ -0,0 +1,108 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command netcdf writes a synthetic climate field to a NetCDF classic
// file, reads it back, and computes the zonal statistics a climate
// workflow starts from. The point is the round trip: dimensions,
// attributes and values survive the file exactly.
//
// Usage: go run ./examples/netcdf
package main
import (
"fmt"
"log"
"math"
"os"
"path/filepath"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
nLat = 36 // 5-degree grid
nLon = 72
)
// A warm anomaly centred on 45 N, 15 E over a zonal gradient, the
// shape of a heat island in a coarse climate model.
field := make([]float64, nLat*nLon)
for i := range nLat {
lat := -90 + 5*(float64(i)+0.5)
for j := range nLon {
lon := -180 + 5*(float64(j)+0.5)
base := 30*math.Cos(lat*math.Pi/180) - 5
dLat := (lat - 45) / 15
dLon := math.Sin((lon - 15) * math.Pi / 180)
anomaly := 8 * math.Exp(-(dLat*dLat + dLon*dLon))
field[i*nLon+j] = base + anomaly
}
}
temp, err := tensor.FromFloats(field, nLat, nLon)
if err != nil {
log.Fatal(err)
}
path := filepath.Join(os.TempDir(), "tensor-example-climate.nc")
defer os.Remove(path)
dims := []tensor.NetCDFDim{
{Name: "lat", Length: nLat},
{Name: "lon", Length: nLon},
}
vars := []tensor.NetCDFVar{{
Name: "temperature",
Dims: []string{"lat", "lon"},
Values: temp,
Attrs: map[string]string{
"units": "degC",
"long_name": "synthetic air temperature",
"anomaly_lon": "15",
},
}}
attrs := map[string]string{
"title": "tensor NetCDF example",
"source": "synthetic Gaussian anomaly",
}
if err := tensor.SaveNetCDF(path, dims, vars, attrs); err != nil {
log.Fatal(err)
}
fmt.Printf("wrote %s: %d x %d grid, %d variables\n\n", path, nLat, nLon, len(vars))
gotDims, gotVars, gotAttrs, err := tensor.LoadNetCDF(path)
if err != nil {
log.Fatal(err)
}
fmt.Printf("dimensions: ")
for _, d := range gotDims {
fmt.Printf("%s(%d) ", d.Name, d.Length)
}
fmt.Printf("\nglobal attributes: %v\n\n", gotAttrs)
back := gotVars[0].Values
if back.Shape()[0] != nLat || back.Shape()[1] != nLon {
log.Fatalf("round trip changed the shape: %v", back.Shape())
}
maxDiff := 0.0
for i := range nLat * nLon {
d := math.Abs(back.FloatAt(i) - field[i])
if d > maxDiff {
maxDiff = d
}
}
fmt.Printf("largest round-trip difference: %g (exact for float64)\n\n", maxDiff)
// Zonal means: the latitude profile of the field, the first thing
// a climate diagnostic asks for.
fmt.Println("latitude zonal mean temperature")
for i := 0; i < nLat; i += 6 {
row, err := tensor.Slice(back, 0, i, i+1)
if err != nil {
log.Fatal(err)
}
mean, err := tensor.Mean(row)
if err != nil {
log.Fatal(err)
}
fmt.Printf("%6.1f° %10.3f degC\n", -90+5*(float64(i)+0.5), mean)
}
}
+135
View File
@@ -0,0 +1,135 @@
// 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")
}
}
+103
View File
@@ -0,0 +1,103 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command pde solves the two canonical one-dimensional partial
// differential equations: the heat equation by Crank-Nicolson and the
// wave equation by velocity Verlet. Both start from the same Gaussian
// pulse on a rod, and the diagnostics show diffusion flattening the
// pulse while the wave keeps its shape and travels.
//
// Usage: go run ./examples/pde
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
n = 201
dx = 0.01 // metres, a 2 m rod
k = 1e-3 // thermal diffusivity, m^2/s
c = 0.5 // wave speed, m/s
)
pulse := make([]float64, n)
for i := range n {
x := float64(i)*dx - 1.0
pulse[i] = math.Exp(-(x * x) / 0.01)
}
u0, err := tensor.FromFloats(pulse, n)
if err != nil {
log.Fatal(err)
}
fmt.Println("heat equation (Crank-Nicolson, ends held at zero):")
fmt.Println(" time peak mean (interior)")
for _, tf := range []float64{0, 0.5, 2, 5} {
var row *tensor.Array
if tf == 0 {
row = u0 // the initial pulse itself
} else {
u, err := tensor.IntegrateHeat1D(u0, k, dx, tf, tf/400, 5, 0, 0)
if err != nil {
log.Fatal(err)
}
// The last sample row holds the final state. The ends are
// Dirichlet zeros, so heat drains out of the rod once the
// pulse reaches them; by the last time printed it has not,
// which is why the interior mean barely moves.
row, err = tensor.Slice(u, 0, 4, 5)
if err != nil {
log.Fatal(err)
}
}
mx, err := tensor.Max(row)
if err != nil {
log.Fatal(err)
}
mean, err := tensor.Mean(row)
if err != nil {
log.Fatal(err)
}
fmt.Printf(" %.2f s %7.4f %7.4f\n", tf, mx.Float(), mean)
}
fmt.Println()
fmt.Println("wave equation (velocity Verlet, fixed ends):")
fmt.Println(" time peak position of the peak")
rest, err := tensor.Zeros(tensor.Float, n)
if err != nil {
log.Fatal(err)
}
for _, tf := range []float64{0, 0.5, 1.0, 1.5} {
var row *tensor.Array
if tf == 0 {
row = u0
} else {
u, err := tensor.IntegrateWave1D(u0, rest, c, dx, tf, tf/600, 5)
if err != nil {
log.Fatal(err)
}
row, err = tensor.Slice(u, 0, 4, 5)
if err != nil {
log.Fatal(err)
}
}
mx, err := tensor.Max(row)
if err != nil {
log.Fatal(err)
}
peak, peakVal := 0, -1.0
for i := range n {
if v := row.FloatAt(i); v > peakVal {
peak, peakVal = i, v
}
}
fmt.Printf(" %.2f s %6.4f x = %.2f m\n",
tf, mx.Float(), float64(peak)*dx-1.0)
}
}
+72
View File
@@ -0,0 +1,72 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command pendulum computes the exact period of a simple pendulum at
// large amplitude through the complete elliptic integral of the first
// kind, and shows how far the small-angle formula drifts once the
// release angle stops being small. The period is
//
// T = 4·sqrt(L/g)·K(sin²(θ₀/2)),
//
// where K is EllipticK with the m = k² parameter convention.
//
// Usage: go run ./examples/pendulum
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
// kComplete evaluates EllipticK at a single parameter.
func kComplete(m float64) float64 {
arr, err := tensor.FromFloats([]float64{m}, 1)
if err != nil {
log.Fatal(err)
}
k, err := tensor.EllipticK(arr)
if err != nil {
log.Fatal(err)
}
v, _ := tensor.FloatAt(k, 0)
return v
}
func main() {
const (
length = 1.0 // metres
grav = 9.80665
)
small := 2 * math.Pi * math.Sqrt(length/grav)
fmt.Println("release angle exact period small-angle period drift")
for _, deg := range []float64{5, 15, 30, 45, 60, 90, 120, 170} {
theta := deg * math.Pi / 180
m := math.Sin(theta/2) * math.Sin(theta/2)
period := 4 * math.Sqrt(length/grav) * kComplete(m)
drift := (period/small - 1) * 100
fmt.Printf("%10.0f° %12.6f s %14.6f s %+6.2f %%\n",
deg, period, small, drift)
}
// The inverse problem: which release angle doubles the small-angle
// period? Bisection on the angle, the period being monotone in it.
target := 2 * small
lo, hi := 0.0, math.Pi
angle := 0.0
for range 80 {
mid := (lo + hi) / 2
m := math.Sin(mid/2) * math.Sin(mid/2)
if 4*math.Sqrt(length/grav)*kComplete(m) < target {
lo = mid
} else {
hi = mid
}
angle = mid
}
fmt.Printf("\na release angle of %.2f° doubles the period (%.4f s)\n",
angle*180/math.Pi, target)
}
+77
View File
@@ -0,0 +1,77 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command qmc compares quasi-random integration against plain Monte
// Carlo on the same two-dimensional integral. Sobol points are a
// digital lattice: every block of 2^m points stratifies each
// coordinate exactly, so the error decays far faster than the
// 1/sqrt(n) of random sampling, and Halton sits in between.
//
// Usage: go run ./examples/qmc
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
// f is the integrand: smooth, with its curvature spread over the unit
// square. The exact value is (1-e^-1)*sqrt(pi)/2*erf(1), the product
// of the x integral and the error function integral over y.
func f(x, y float64) float64 { return math.Exp(-x - y*y) }
const exact = 0.4720828881800443 // (1-e^-1)*sqrt(pi)/2*erf(1)
// estimate integrates f over [0,1]^2 from an (n,2) point set.
func estimate(pts *tensor.Array, n int) float64 {
s := 0.0
for i := range n {
x, err := tensor.FloatAt(pts, i, 0)
if err != nil {
log.Fatal(err)
}
y, err := tensor.FloatAt(pts, i, 1)
if err != nil {
log.Fatal(err)
}
s += f(x, y)
}
return s / float64(n)
}
func main() {
fmt.Printf("integral of exp(-x - y^2) over the unit square, exact %.10f\n\n", exact)
fmt.Println(" points Monte Carlo Halton Sobol")
for _, n := range []int{64, 256, 1024, 4096, 16384} {
// Monte Carlo: uniform draws from the seeded generator.
g := tensor.NewGenerator(int64(n))
mc, err := tensor.Floats(g, 2*n)
if err != nil {
log.Fatal(err)
}
mcPts, err := tensor.Reshape(mc, n, 2)
if err != nil {
log.Fatal(err)
}
// Halton and Sobol from the first point on; Sobol skips its
// origin point exactly as Halton does.
hal, err := tensor.HaltonPoints(n, 2, 0)
if err != nil {
log.Fatal(err)
}
sob, err := tensor.SobolPoints(n, 2, 0)
if err != nil {
log.Fatal(err)
}
eMC := math.Abs(estimate(mcPts, n) - exact)
eHal := math.Abs(estimate(hal, n) - exact)
eSob := math.Abs(estimate(sob, n) - exact)
fmt.Printf(" %6d %.3e %.3e %.3e\n", n, eMC, eHal, eSob)
}
fmt.Println()
fmt.Println("the quasi-random errors collapse with n; the Monte Carlo")
fmt.Println("error only shrinks as 1/sqrt(n) and stays noisy on top")
}
+118
View File
@@ -0,0 +1,118 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command regression fits a linear trend to a noisy time series,
// reports the full inference table (coefficients, standard errors,
// t-statistics, p-values, R²) and checks that the residuals are
// actually uncorrelated, which is the assumption the t-tests rest on.
//
// Usage: go run ./examples/regression
package main
import (
"fmt"
"log"
"sourcedock.dev/petrbalvin/tensor"
"sourcedock.dev/petrbalvin/tensor/signal"
"sourcedock.dev/petrbalvin/tensor/stats"
)
func main() {
const n = 400
// A trend of 0.05 per sample on a level of 2, with AR(1) noise
// (rho = 0.3), drawn from the reproducible generator.
g := tensor.NewGenerator(7)
white, err := tensor.Normal(g, n, 0, 1)
if err != nil {
log.Fatal(err)
}
y := make([]float64, n)
ar := 0.0
for i := range n {
w, _ := tensor.FloatAt(white, i)
ar = 0.3*ar + w
y[i] = 2 + 0.05*float64(i) + 0.4*ar
}
yArr, err := tensor.FromFloats(y, n)
if err != nil {
log.Fatal(err)
}
// The design carries its own intercept column, the convention of
// the classic linear model.
design := make([]float64, 2*n)
for i := range n {
design[2*i] = 1
design[2*i+1] = float64(i)
}
xArr, err := tensor.FromFloats(design, n, 2)
if err != nil {
log.Fatal(err)
}
fit, err := stats.LinearRegression(xArr, yArr)
if err != nil {
log.Fatal(err)
}
fmt.Println("ordinary least squares fit, y = intercept + slope * t")
fmt.Println("term estimate std error t-stat p-value")
fmt.Printf("intercept %9.4f %9.4f %7.3f %.3g\n",
fit.Coefficients[0], fit.StandardErrors[0], fit.TStatistics[0], fit.PValues[0])
fmt.Printf("slope %9.4f %9.4f %7.3f %.3g\n",
fit.Coefficients[1], fit.StandardErrors[1], fit.TStatistics[1], fit.PValues[1])
fmt.Printf("\nR² = %.4f, adjusted R² = %.4f, residual variance = %.4f\n",
fit.RSquared, fit.AdjustedRSquared, fit.ResidualVariance)
fmt.Println("(the generating values were intercept 2, slope 0.05)")
// The t-tests assume uncorrelated residuals. Pull them out and
// check the autocorrelation at the first few lags; with rho = 0.3
// in the noise, lag 1 must show clear correlation, which is the
// honest caveat for the standard errors above.
resid := make([]float64, n)
for i := range n {
pred := fit.Coefficients[0] + fit.Coefficients[1]*float64(i)
resid[i] = y[i] - pred
}
rArr, err := tensor.FromFloats(resid, n)
if err != nil {
log.Fatal(err)
}
ac, err := signal.Autocorrelate(rArr, 5)
if err != nil {
log.Fatal(err)
}
// The transform returns lags 0..5; lag 0 is 1 by definition, the
// AR(1) memory shows from lag 1 on.
fmt.Print("\nresidual autocorrelation:")
for lag := 1; lag <= 5; lag++ {
v, _ := tensor.FloatAt(ac, lag)
fmt.Printf(" lag %d: %+.3f", lag, v)
}
fmt.Println()
// A two-sample test on the first and last halves: with a trend of
// 0.05 over 200 samples the means must differ decisively.
first, err := tensor.Slice(yArr, 0, 0, n/2)
if err != nil {
log.Fatal(err)
}
last, err := tensor.Slice(yArr, 0, n/2, n)
if err != nil {
log.Fatal(err)
}
t, df, p, err := stats.WelchTTest(first, last)
if err != nil {
log.Fatal(err)
}
meanOf := func(a *tensor.Array) float64 {
m, err := tensor.Mean(a)
if err != nil {
log.Fatal(err)
}
return m
}
fmt.Printf("\nWelch t-test, first half vs second half:\n")
fmt.Printf(" means %.3f vs %.3f, t = %.2f, df = %.1f, p = %.3g\n",
meanOf(first), meanOf(last), t, df, p)
}
+136
View File
@@ -0,0 +1,136 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command spectral estimates the frequency content of a signal two
// ways: Welch's averaged periodogram on evenly sampled data, and the
// Lomb-Scargle periodogram on the same signal observed at irregular
// times, where an FFT cannot run at all. Both must find the two
// buried sinusoids.
//
// Usage: go run ./examples/spectral
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
fs = 100.0
seconds = 4.0
f1 = 5.0
f2 = 13.0
)
n := int(fs * seconds)
gen := tensor.NewGenerator(11)
// The signal: two sinusoids plus noise.
t := make([]float64, n)
x := make([]float64, n)
for i := range n {
t[i] = float64(i) / fs
x[i] = math.Sin(2*math.Pi*f1*t[i]) + 0.6*math.Sin(2*math.Pi*f2*t[i]) + 0.4*gen.NormalUnit()
}
xArr, err := tensor.FromFloats(x, n)
if err != nil {
log.Fatal(err)
}
// Welch: average periodograms over Hann-windowed segments, the
// variance-suppressed estimate an FFT alone cannot give.
freqs, psd, err := tensor.WelchPSD(xArr, fs, 256, 128, "hann")
if err != nil {
log.Fatal(err)
}
wf1, wf2 := twoPeaks(peakFrequencies(freqs, psd, 2))
fmt.Printf("welch peaks at %.2f Hz and %.2f Hz (want %.1f and %.1f)\n", wf1, wf2, f1, f2)
// Lomb-Scargle: keep every second sample at jittered times, the
// uneven regime the DFT does not define. The mean rate stays at
// 50 Hz, comfortably above both sources' Nyquist needs, while the
// jitter is what makes the ordinary FFT inapplicable.
times := make([]float64, 0, n/2)
values := make([]float64, 0, n/2)
for i := 0; i < n; i += 2 {
jitter := 0.6 * gen.Unit() / fs
times = append(times, t[i]+jitter)
values = append(values, x[i])
}
tArr, err := tensor.FromFloats(times, len(times))
if err != nil {
log.Fatal(err)
}
vArr, err := tensor.FromFloats(values, len(values))
if err != nil {
log.Fatal(err)
}
lsFreqs, power, err := tensor.LombScargle(tArr, vArr, 1.0, 30.0, 3000)
if err != nil {
log.Fatal(err)
}
lf1, lf2 := twoPeaks(peakFrequencies(lsFreqs, power, 2))
fmt.Printf("lomb-scargle peaks at %.2f Hz and %.2f Hz (want %.1f and %.1f)\n", lf1, lf2, f1, f2)
for _, got := range []float64{wf1, wf2, lf1, lf2} {
if math.Abs(got-f1) > 0.3 && math.Abs(got-f2) > 0.3 {
log.Fatalf("a peak landed at %.2f Hz, away from both sources", got)
}
}
}
// peakFrequencies returns the abscissae of the count largest local
// maxima of a periodogram, descending by height and kept at least
// 1.5 Hz apart so a sidelobe of a tall peak cannot shadow a real one.
func peakFrequencies(freqs, power *tensor.Array, count int) []float64 {
n := freqs.Len()
// Three-point boxcar smooth: the periodogram's noise is white, a
// genuine peak is not.
smooth := make([]float64, n)
for i := range n {
lo := max(i-1, 0)
hi := min(i+1, n-1)
s := 0.0
for j := lo; j <= hi; j++ {
s += power.FloatAt(j)
}
smooth[i] = s / float64(hi-lo+1)
}
type peak struct {
f, h float64
}
var peaks []peak
for i := 1; i < n-1; i++ {
if smooth[i] > smooth[i-1] && smooth[i] >= smooth[i+1] {
peaks = append(peaks, peak{freqs.FloatAt(i), smooth[i]})
}
}
for i := 1; i < len(peaks); i++ {
for j := i; j > 0 && peaks[j-1].h < peaks[j].h; j-- {
peaks[j-1], peaks[j] = peaks[j], peaks[j-1]
}
}
out := make([]float64, 0, count)
for _, p := range peaks {
if len(out) == count {
break
}
far := true
for _, f := range out {
if math.Abs(p.f-f) < 1.5 {
far = false
break
}
}
if far {
out = append(out, p.f)
}
}
return out
}
// twoPeaks unpacks the two-element result of peakFrequencies.
func twoPeaks(fs []float64) (float64, float64) { return fs[0], fs[1] }
+139
View File
@@ -0,0 +1,139 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command wavelets demonstrates the discrete wavelet transform on a
// denoising task and the continuous transform on a time-frequency
// task: a clean signal is buried in noise, the detail coefficients are
// soft-thresholded and the signal rebuilt, then a two-tone signal with
// an abrupt frequency change is mapped by the CWT so the change is
// visible in time, not just in frequency.
//
// Usage: go run ./examples/wavelets
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
"sourcedock.dev/petrbalvin/tensor/signal"
)
func main() {
const n = 1024
// A clean decaying sinusoid, buried in noise drawn from the
// reproducible generator so the run is exactly repeatable.
g := tensor.NewGenerator(2026)
noise, err := tensor.Normal(g, n, 0, 0.25)
if err != nil {
log.Fatal(err)
}
clean := make([]float64, n)
dirty := make([]float64, n)
for i := range n {
x := float64(i) / n
clean[i] = math.Sin(2*math.Pi*3*x) * math.Exp(-3*x)
nv, _ := tensor.FloatAt(noise, i)
dirty[i] = clean[i] + nv
}
dirtyArr, err := tensor.FromFloats(dirty, n)
if err != nil {
log.Fatal(err)
}
// Decompose, soft-threshold the detail coefficients, rebuild. The
// threshold sits at twice the noise standard deviation, the level
// where a noise-only coefficient almost never survives.
const levels = 5
coef, err := signal.DWT(dirtyArr, levels)
if err != nil {
log.Fatal(err)
}
approx := n >> levels
const threshold = 2 * 0.25
raw := coef.RawFloats()
for i := approx; i < len(raw); i++ {
v := raw[i]
switch {
case v > threshold:
raw[i] = v - threshold
case v < -threshold:
raw[i] = v + threshold
default:
raw[i] = 0
}
}
denoised, err := signal.IDWT(coef, levels)
if err != nil {
log.Fatal(err)
}
mse := func(a []float64) float64 {
s := 0.0
for i := range n {
d := a[i] - clean[i]
s += d * d
}
return s / float64(n)
}
fmt.Println("mean squared error against the clean signal:")
fmt.Printf(" noisy %.6f\n", mse(dirty))
fmt.Printf(" denoised %.6f\n", mse(denoised.RawFloats()[:n]))
fmt.Println()
// The continuous transform: 512 samples of a signal whose tone
// jumps from 8 to 32 cycles over the whole run, halfway through.
// A Morlet scale a responds at omega0/(2*pi*a) cycles per sample,
// which is omega0*N/(2*pi*a) cycles per record of N = 512 samples,
// so with omega0 = 5 the two tones live near a = 51 and a = 13;
// the scalogram ridge must jump between them.
const m = 512
chirp := make([]float64, m)
for i := range m {
freq := 8.0
if i >= m/2 {
freq = 32.0
}
chirp[i] = math.Sin(2 * math.Pi * freq * float64(i) / m)
}
chirpArr, err := tensor.FromFloats(chirp, m)
if err != nil {
log.Fatal(err)
}
scales := []float64{4, 8, 13, 16, 26, 32, 51, 64}
scalogram, err := signal.CWT(chirpArr, signal.Morlet, scales, 1)
if err != nil {
log.Fatal(err)
}
fmt.Println("CWT ridge: the scale carrying the peak energy in each half")
// The wavelet of scale 64 spans about 256 samples, so the outer
// quarters of the run are edge territory; the ridge is read from
// the interior of each half only.
const margin = 128
for _, seg := range []struct {
label string
start, stop int
}{
{"first half ", margin, m/2 - margin/2},
{"second half", m/2 + margin/2, m - margin},
} {
best := 0
bestMag := -1.0
for si := range scales {
for i := seg.start; i < seg.stop; i++ {
// The scalogram is (len(scales), m), one complex row
// per scale; the ridge is the peak magnitude.
cv, err := tensor.ComplexAt(scalogram, si, i)
if err != nil {
log.Fatal(err)
}
if a := math.Hypot(real(cv), imag(cv)); a > bestMag {
best, bestMag = si, a
}
}
}
fmt.Printf(" %s: scale %.0f\n", seg.label, scales[best])
}
}
+785
View File
@@ -0,0 +1,785 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package tensor is the public facade: the overview and the guarantees
// live in doc.go, and this file only forwards.
package tensor
import (
grad "sourcedock.dev/petrbalvin/tensor/grad"
integrate "sourcedock.dev/petrbalvin/tensor/integrate"
core "sourcedock.dev/petrbalvin/tensor/internal/core"
io "sourcedock.dev/petrbalvin/tensor/io"
linalg "sourcedock.dev/petrbalvin/tensor/linalg"
optim "sourcedock.dev/petrbalvin/tensor/optim"
plot "sourcedock.dev/petrbalvin/tensor/plot"
signal "sourcedock.dev/petrbalvin/tensor/signal"
stats "sourcedock.dev/petrbalvin/tensor/stats"
)
// Wavelet names for signal.CWT.
const (
Morlet = signal.Morlet
MexicanHat = signal.MexicanHat
)
// Core types.
type Array = core.Array
type Dtype = core.Dtype
type Generator = core.Generator
type JacobianOptions = core.JacobianOptions
type Scalar = core.Scalar
type SparseCOO = core.SparseCOO
type BoundaryConditions = integrate.BoundaryConditions
type BDFVarOptions = integrate.BDFVarOptions
type BDFVarStats = integrate.BDFVarStats
type CollocationOptions = integrate.CollocationOptions
type CollocationSolution = integrate.CollocationSolution
type CubicSpline = linalg.CubicSpline
type DAEOptions = integrate.DAEOptions
type FITSTable = io.FITSTable
type FITSTableColumn = io.FITSTableColumn
type HDF5Dataset = io.HDF5Dataset
type HDF5TextDataset = io.HDF5TextDataset
type HDF5WriteOptions = io.HDF5WriteOptions
type NetCDFDim = io.NetCDFDim
type NetCDFVar = io.NetCDFVar
type HMCOptions = grad.HMCOptions
type NewtonCGOptions = grad.NewtonCGOptions
type HessianOptions = grad.HessianOptions
type KalmanOptions = signal.KalmanOptions
type KalmanResult = signal.KalmanResult
type StateFunc = signal.StateFunc
type JacobianFunc = signal.JacobianFunc
type ARMAOptions = signal.ARMAOptions
type ARMAResult = signal.ARMAResult
type ElasticNetResult = stats.ElasticNetResult
type LassoPathResult = stats.LassoPathResult
type HuberRegressionResult = stats.HuberRegressionResult
type QuantileRegressionResult = stats.QuantileRegressionResult
type PCAResult = stats.PCAResult
type KMeansResult = stats.KMeansResult
type GaussianMixtureResult = stats.GaussianMixtureResult
type LinearMixedModelResult = stats.LinearMixedModelResult
type HiddenMarkovModel = stats.HiddenMarkovModel
type HiddenMarkovFitResult = stats.HiddenMarkovFitResult
type Dendrogram = stats.Dendrogram
type Alternative = stats.Alternative
type Linkage = stats.Linkage
type Kernel = stats.Kernel
type GaussianProcessResult = stats.GaussianProcessResult
type LinearProgramOptions = optim.LinearProgramOptions
type QPOptions = optim.QPOptions
type CMAESOptions = optim.CMAESOptions
type SimulatedAnnealingOptions = optim.SimulatedAnnealingOptions
type NonlinearConstraints = optim.NonlinearConstraints
type CubatureOptions = integrate.CubatureOptions
type TriangleMesh2D = integrate.TriangleMesh2D
type TetraMesh3D = integrate.TetraMesh3D
type FEMPoissonOptions = integrate.FEMPoissonOptions
type FEMPoisson3DOptions = integrate.FEMPoisson3DOptions
type MidpointOptions = integrate.MidpointOptions
type FilonOptions = integrate.FilonOptions
type LeastSquaresInfo = linalg.LeastSquaresInfo
type LinearRegressionResult = stats.LinearRegressionResult
type LogisticRegressionResult = stats.LogisticRegressionResult
type PoissonRegressionResult = stats.PoissonRegressionResult
type BrentOptions = optim.BrentOptions
type CWTWavelet = signal.CWTWavelet
type Daubechies = signal.Daubechies
type DifferentialEvolutionOptions = optim.DifferentialEvolutionOptions
type LBFGSOptions = optim.LBFGSOptions
type LinearConstraints = optim.LinearConstraints
type LMOptions = optim.LMOptions
type FitResult = optim.FitResult
type FitStatus = optim.FitStatus
type DWTMode = signal.DWTMode
type MinimiseOptions = optim.MinimiseOptions
type ODEEventHit = integrate.ODEEventHit
type ODEOptions = integrate.ODEOptions
type ODEWatch = integrate.ODEWatch
type Pipeline = linalg.Pipeline
type QuadratureOptions = integrate.QuadratureOptions
type RootSystemOptions = optim.RootSystemOptions
type SparseCSR = linalg.SparseCSR
type SparseCSC = linalg.SparseCSC
type SparseCholesky = linalg.SparseCholesky
type SparseLU = linalg.SparseLU
type SparseOrdering = linalg.SparseOrdering
type SparseILU = linalg.SparseILU
type STFTOptions = signal.STFTOptions
type Tensor = grad.Tensor
// Charts.
type Chart = plot.Chart
type Point = plot.Point
type Series = plot.Series
// Core constants.
const (
Bool = core.Bool
Float = core.Float
Float16 = core.Float16
Float32 = core.Float32
Int = core.Int
Int8 = core.Int8
Uint8 = core.Uint8
Int16 = core.Int16
Uint16 = core.Uint16
Int32 = core.Int32
Uint32 = core.Uint32
Complex = core.Complex
)
// Sparse factorisation orderings.
const (
SparseOrderingNatural = linalg.SparseOrderingNatural
SparseOrderingReverseCuthillMcKee = linalg.SparseOrderingReverseCuthillMcKee
SparseOrderingMinimumDegree = linalg.SparseOrderingMinimumDegree
)
// Wavelet extension modes.
const (
DWTPeriodic = signal.DWTPeriodic
DWTZeroPad = signal.DWTZeroPad
)
// Daubechies wavelet families.
const (
DB2 = signal.DB2
DB3 = signal.DB3
DB4 = signal.DB4
DB5 = signal.DB5
DB6 = signal.DB6
DB7 = signal.DB7
DB8 = signal.DB8
)
// Least-squares stopping criteria.
const (
LeastSquaresResidual = linalg.LeastSquaresResidual
LeastSquaresNormal = linalg.LeastSquaresNormal
LeastSquaresCondition = linalg.LeastSquaresCondition
)
// Alternatives of the directional tests.
const (
TwoSided = stats.TwoSided
Less = stats.Less
Greater = stats.Greater
)
// Hierarchical clustering linkages.
const (
SingleLinkage = stats.SingleLinkage
CompleteLinkage = stats.CompleteLinkage
AverageLinkage = stats.AverageLinkage
CentroidLinkage = stats.CentroidLinkage
WardLinkage = stats.WardLinkage
)
// Fit statuses.
const (
FitConverged = optim.FitConverged
FitStalled = optim.FitStalled
FitBudget = optim.FitBudget
)
// Robust regression constants.
const (
DefaultHuberTuning = stats.DefaultHuberTuning
TheilSenMaxObservations = stats.TheilSenMaxObservations
)
var ArgMaxAxis = core.ArgMaxAxis
var FromArray = grad.FromArray
var FromFloat64s = grad.FromFloat64s
var ArgMinAxis = core.ArgMinAxis
var TopK = core.TopK
var FromComplexes = core.FromComplexes
var FromFloat32s = core.FromFloat32s
var Abs = core.Abs
var Add = core.Add
var AddC = core.AddC
var AddF = core.AddF
var AddI = core.AddI
var Airy = core.Airy
var AnalyticSignal = signal.AnalyticSignal
var ArgSort = core.ArgSort
var Argwhere = core.Argwhere
var ARMASpectrum = signal.ARMASpectrum
var AssignBins = core.AssignBins
var Astype = core.Astype
var BesselI0 = core.BesselI0
var BesselI1 = core.BesselI1
var BesselIn = core.BesselIn
var BesselK0 = core.BesselK0
var BesselK1 = core.BesselK1
var BesselKn = core.BesselKn
var And = core.And
var Beta = core.Beta
var BoolsFromArray = core.BoolsFromArray
var BroadcastTo = core.BroadcastTo
var Ceil = core.Ceil
var ChebyshevT = core.ChebyshevT
var ChebyshevU = core.ChebyshevU
var Chirp = signal.Chirp
var ChiSquareIndependence = stats.ChiSquareIndependence
var ClipF = core.ClipF
var ClipI = core.ClipI
var Col = core.Col
var ComplexFromArray = core.ComplexFromArray
var Concat = core.Concat
var Copy = core.Copy
var CramersV = stats.CramersV
var Cos = core.Cos
var Cosm1 = core.Cosm1
var CrossProduct = core.CrossProduct
var CumProd = core.CumProd
var CumSum = core.CumSum
var CumulativeIntegrate = core.CumulativeIntegrate
var Diag = core.Diag
var Diagonal = core.Diagonal
var Diff = core.Diff
var Digamma = core.Digamma
var DirichletDensity = stats.DirichletDensity
var DirichletDraws = stats.DirichletDraws
var DirichletMean = stats.DirichletMean
var DirichletMode = stats.DirichletMode
var Div = core.Div
var DivC = core.DivC
var DivF = core.DivF
var DivI = core.DivI
var Dot = core.Dot
var Einsum = core.Einsum
var Eq = core.Eq
var EqF = core.EqF
var EqI = core.EqI
var ElasticNet = stats.ElasticNet
var Erf = core.Erf
var Erfc = core.Erfc
var Envelope = signal.Envelope
var EstimateAR = signal.EstimateAR
var EstimateARMA = signal.EstimateARMA
var EvaluatePolynomial = core.EvaluatePolynomial
var Exp = core.Exp
var ExpIntegralE1 = core.ExpIntegralE1
var ExpIntegralEi = core.ExpIntegralEi
var ExtendedKalmanFilter = signal.ExtendedKalmanFilter
var Filtfilt = signal.Filtfilt
var FisherExactTest = stats.FisherExactTest
var FitHiddenMarkovModel = stats.FitHiddenMarkovModel
var Flatten = core.Flatten
var Flip = core.Flip
var Float32s = core.Float32s
var Floats = core.Floats
var FloatsFromArray = core.FloatsFromArray
var Floor = core.Floor
var FresnelC = core.FresnelC
var FresnelS = core.FresnelS
var FromBools = core.FromBools
var FromBytes = core.FromBytes
var FromFloat32Slice = core.FromFloat32Slice
var FromFloat16s = core.FromFloat16s
var FromFloatSlice = core.FromFloatSlice
var FromFloats = core.FromFloats
var FromInt16s = core.FromInt16s
var FromInt32s = core.FromInt32s
var FromInt8s = core.FromInt8s
var FromInts = core.FromInts
var FromUint16s = core.FromUint16s
var FromUint32s = core.FromUint32s
var FromUint8s = core.FromUint8s
var FullC = core.FullC
var FullF = core.FullF
var FullF16 = core.FullF16
var FullF32s = core.FullF32s
var FullI = core.FullI
var HalvesFromArray = core.HalvesFromArray
var HalfFromFloat64 = core.HalfFromFloat64
var HalfToFloat64 = core.HalfToFloat64
var GaussianMixture = stats.GaussianMixture
var GaussianMixtureBIC = stats.GaussianMixtureBIC
var GaussianProcessRegression = stats.GaussianProcessRegression
var Gamma = core.Gamma
var Gather = core.Gather
var Ge = core.Ge
var GeF = core.GeF
var GeI = core.GeI
var Grid = core.Grid
var Gt = core.Gt
var GtF = core.GtF
var GtI = core.GtI
var Hermite = core.Hermite
var HierarchicalClustering = stats.HierarchicalClustering
var Identity = core.Identity
var Int16sFromArray = core.Int16sFromArray
var Int32sFromArray = core.Int32sFromArray
var Int8sFromArray = core.Int8sFromArray
var Interpolate = core.Interpolate
var Interpolate2D = core.Interpolate2D
var InterpolateGrid = core.InterpolateGrid
var InterpolateMonotone = core.InterpolateMonotone
var Ints = core.Ints
var IntsFromArray = core.IntsFromArray
var IsFinite = core.IsFinite
var IsInf = core.IsInf
var IsNaN = core.IsNaN
var Jacobian = core.Jacobian
var Kron = core.Kron
var Laguerre = core.Laguerre
var Le = core.Le
var LeF = core.LeF
var LeI = core.LeI
var Legendre = core.Legendre
var LegendreAssociated = core.LegendreAssociated
var Linspace = core.Linspace
var LinearMixedModel = stats.LinearMixedModel
var LnGamma = core.LnGamma
var Log = core.Log
var Log10 = core.Log10
var Log2 = core.Log2
var LognormalCDF = stats.LognormalCDF
var LognormalDensity = stats.LognormalDensity
var LognormalQuantile = stats.LognormalQuantile
var LowerTriangle = core.LowerTriangle
var Lt = core.Lt
var LtF = core.LtF
var LtI = core.LtI
var MatMul2D = core.MatMul2D
var Max = core.Max
var MaxAxis = core.MaxAxis
var Maximum = core.Maximum
var McNemarTest = stats.McNemarTest
var MeanAxis = core.MeanAxis
var Min = core.Min
var MinAxis = core.MinAxis
var Minimum = core.Minimum
var MoveAxis = core.MoveAxis
var Mul = core.Mul
var MulC = core.MulC
var MulF = core.MulF
var MulI = core.MulI
var Ne = core.Ne
var NeF = core.NeF
var NeI = core.NeI
var New = core.New
var NewGenerator = core.NewGenerator
var NewHiddenMarkovModel = stats.NewHiddenMarkovModel
var NewSparseCOO = core.NewSparseCOO
var NewTetraMesh3D = integrate.NewTetraMesh3D
var Norm = core.Norm
var Normal = core.Normal
var Not = core.Not
var OneHot = core.OneHot
var Ones = core.Ones
var OnesLike = core.OnesLike
var Or = core.Or
var Pad = core.Pad
var Permutation = core.Permutation
var Pow = core.Pow
var PowI = core.PowI
var Prod = core.Prod
var Quo = core.Quo
var QuoI = core.QuoI
var Range = core.Range
var RangeBy = core.RangeBy
var Repeat = core.Repeat
var Reshape = core.Reshape
var Resample = signal.Resample
var ResampleFourier = signal.ResampleFourier
var Reverse = core.Reverse
var RRQR = linalg.RRQR
var RRQRRank = linalg.RRQRRank
var Roll = core.Roll
var RollingMax = stats.RollingMax
var RollingMean = stats.RollingMean
var RollingMin = stats.RollingMin
var RollingSum = stats.RollingSum
var Round = core.Round
var Row = core.Row
var Scatter = core.Scatter
var SearchSorted = core.SearchSorted
var Select = core.Select
var Shuffle = core.Shuffle
var Splitmix64 = core.Splitmix64
var Substream = core.Substream
var Sigmoid = core.Sigmoid
var Sign = core.Sign
var Sin = core.Sin
var Sinc = core.Sinc
var Slice = core.Slice
var Sort = core.Sort
var SpAdd = core.SpAdd
var SpMatMul = core.SpMatMul
var SpMul = core.SpMul
var SparseFrom = core.SparseFrom
var SphericalBesselJ = core.SphericalBesselJ
var SphericalBesselY = core.SphericalBesselY
var SphericalHarmonic = core.SphericalHarmonic
var Spectrogram = signal.Spectrogram
var STFT = signal.STFT
var SphericalHarmonicReal = core.SphericalHarmonicReal
var Sqrt = core.Sqrt
var Squeeze = core.Squeeze
var Stack = core.Stack
var Sub = core.Sub
var SubC = core.SubC
var SubF = core.SubF
var SubI = core.SubI
var Sum = core.Sum
var SumAxis = core.SumAxis
var Take = core.Take
var Tan = core.Tan
var Tanh = core.Tanh
var Tile = core.Tile
var Transpose = core.Transpose
var TransposeAxes = core.TransposeAxes
var TrimmedMean = stats.TrimmedMean
var Trigamma = core.Trigamma
var Trunc = core.Trunc
var TruncatedNormal = core.TruncatedNormal
var Uint16sFromArray = core.Uint16sFromArray
var Uint32sFromArray = core.Uint32sFromArray
var Uint8sFromArray = core.Uint8sFromArray
var Unique = core.Unique
var Unsqueeze = core.Unsqueeze
var UpperTriangle = core.UpperTriangle
var Where = core.Where
var WithComplex = core.WithComplex
var WithFloat = core.WithFloat
var WithInt = core.WithInt
var Xor = core.Xor
var Zeros = core.Zeros
var ZerosLike = core.ZerosLike
// Functions re-exported from the domain packages.
var ANOVAOneWay = stats.ANOVAOneWay
var AdaptiveAvgPool1D = signal.AdaptiveAvgPool1D
var AdaptiveAvgPool2D = signal.AdaptiveAvgPool2D
var AdaptiveAvgPool3D = signal.AdaptiveAvgPool3D
var AdaptiveMaxPool1D = signal.AdaptiveMaxPool1D
var AdaptiveMaxPool2D = signal.AdaptiveMaxPool2D
var AdaptiveMaxPool3D = signal.AdaptiveMaxPool3D
var AdjointODE = grad.AdjointODE
var Hessian = grad.Hessian
var HessianVectorProduct = grad.HessianVectorProduct
var MinimiseNewtonCG = grad.MinimiseNewtonCG
var All = core.All
var Any = core.Any
var ArgMax = core.ArgMax
var ArgMin = core.ArgMin
var ArrayFromFloatsSafe = linalg.ArrayFromFloatsSafe
var AvgPool1D = signal.AvgPool1D
var AvgPool2D = signal.AvgPool2D
var AvgPool3D = signal.AvgPool3D
var BenjaminiHochberg = stats.BenjaminiHochberg
var BesselJ = core.BesselJ
var BesselJRealOrder = core.BesselJRealOrder
var BesselY = core.BesselY
var BetaIncomplete = stats.BetaIncomplete
var BinCounts = stats.BinCounts
var BinomialCDF = stats.BinomialCDF
var BinomialDraws = stats.BinomialDraws
var BinomialQuantile = stats.BinomialQuantile
var Bonferroni = stats.Bonferroni
var BoolAt = core.BoolAt
var BootstrapCI = stats.BootstrapCI
var BoxTetraMesh3D = integrate.BoxTetraMesh3D
var BroadcastWith = core.BroadcastWith
var ChiSquareCDF = stats.ChiSquareCDF
var ChiSquareDraws = stats.ChiSquareDraws
var ChiSquareGoodnessOfFit = stats.ChiSquareGoodnessOfFit
var ChiSquareQuantile = stats.ChiSquareQuantile
var Cholesky = linalg.Cholesky
var CholeskyDowndate = linalg.CholeskyDowndate
var CholeskyUpdate = linalg.CholeskyUpdate
var ComplexAt = core.ComplexAt
var Cond = linalg.Cond
var Conv1D = signal.Conv1D
var Conv2D = signal.Conv2D
var Conv2DGroups = signal.Conv2DGroups
var Conv3D = signal.Conv3D
var ConvTranspose2D = signal.ConvTranspose2D
var Correlation = core.Correlation
var CorrelationMatrix = stats.CorrelationMatrix
var CountNonzero = core.CountNonzero
var Covariance = core.Covariance
var CovarianceMatrix = stats.CovarianceMatrix
var DaubechiesDWT = signal.DaubechiesDWT
var DaubechiesIDWT = signal.DaubechiesIDWT
var DCT = signal.DCT
var DST = signal.DST
var Decimate = signal.Decimate
var Det = linalg.Det
var DetComplex = linalg.DetComplex
var Eigen = linalg.Eigen
var EigenComplex = linalg.EigenComplex
var EigenGeneral = linalg.EigenGeneral
var EigenGeneralised = linalg.EigenGeneralised
var Equal = core.Equal
var ExponentialCDF = stats.ExponentialCDF
var ExponentialDraws = stats.ExponentialDraws
var ExponentialQuantile = stats.ExponentialQuantile
var FFT = signal.FFT
var FFT2 = signal.FFT2
var FFT3 = signal.FFT3
var FFTFreq = signal.FFTFreq
var FFTN = signal.FFTN
var FindRoot = optim.FindRoot
var FindRootNewton = optim.FindRootNewton
var FindRootBrent = optim.FindRootBrent
var FindRootSystem = optim.FindRootSystem
var FitPolynomial = linalg.FitPolynomial
var FloatAt = core.FloatAt
var GMRES = linalg.GMRES
var GammaCDF = stats.GammaCDF
var GammaDraws = stats.GammaDraws
var GammaLower = stats.GammaLower
var GammaQuantile = stats.GammaQuantile
var GammaUpper = stats.GammaUpper
var GaussLegendreNodes = integrate.GaussLegendreNodes
var GlobalAvgPool1D = signal.GlobalAvgPool1D
var GlobalAvgPool2D = signal.GlobalAvgPool2D
var GlobalAvgPool3D = signal.GlobalAvgPool3D
var GlobalMaxPool1D = signal.GlobalMaxPool1D
var GlobalMaxPool2D = signal.GlobalMaxPool2D
var GlobalMaxPool3D = signal.GlobalMaxPool3D
var Gradient1D = signal.Gradient1D
var Histogram = stats.Histogram
var Histogram2D = stats.Histogram2D
var Holm = stats.Holm
var HuberRegression = stats.HuberRegression
var HuberRegressionTuned = stats.HuberRegressionTuned
var IDCT = signal.IDCT
var IDST = signal.IDST
var IFFT = signal.IFFT
var IFFT2 = signal.IFFT2
var IFFT3 = signal.IFFT3
var IFFTN = signal.IFFTN
var IRFFT = signal.IRFFT
var IntAt = core.IntAt
var Integrate = core.Integrate
var IntegrateAdvection1D = integrate.IntegrateAdvection1D
var IntegrateAdvectionDiffusion1D = integrate.IntegrateAdvectionDiffusion1D
var IntegrateBDF2 = integrate.IntegrateBDF2
var IntegrateBDFVar = integrate.IntegrateBDFVar
var IntegrateBackwardEuler = integrate.IntegrateBackwardEuler
var IntegrateBoundary = integrate.IntegrateBoundary
var IntegrateDAE = integrate.IntegrateDAE
var IntegrateFunction = integrate.IntegrateFunction
var IntegrateFilon = integrate.IntegrateFilon
var IntegrateMidpoint = integrate.IntegrateMidpoint
var IntegrateODE = integrate.IntegrateODE
var IntegrateODEPath = integrate.IntegrateODEPath
var IntegrateODESteps = integrate.IntegrateODESteps
var IntegrateROS4 = integrate.IntegrateROS4
var IntegrateRK4 = integrate.IntegrateRK4
var IntegrateUpwindAdvection1D = integrate.IntegrateUpwindAdvection1D
var IntegrateVerlet = integrate.IntegrateVerlet
var IntegrateYoshida4 = integrate.IntegrateYoshida4
var Inv = linalg.Inv
var Item = core.Item
var KalmanFilter = signal.KalmanFilter
var UnscentedKalmanFilter = signal.UnscentedKalmanFilter
var KendallTau = stats.KendallTau
var KolmogorovSmirnovTest = stats.KolmogorovSmirnovTest
var KernelDensity = stats.KernelDensity
var KMeans = stats.KMeans
var Laplacian = signal.Laplacian
var LeastSquares = linalg.LeastSquares
var Lasso = stats.Lasso
var LassoPath = stats.LassoPath
var LevenbergMarquardt = optim.LevenbergMarquardt
var LevenbergMarquardtFit = optim.LevenbergMarquardtFit
var LogisticRegression = stats.LogisticRegression
var LnFactorial = core.LnFactorial
var LoadCSV = io.LoadCSV
var LoadCSVReader = io.LoadCSVReader
var LoadFITS = io.LoadFITS
var LoadHDF5 = io.LoadHDF5
var LoadNetCDF = io.LoadNetCDF
var LombScargle = signal.LombScargle
var MannWhitneyU = stats.MannWhitneyU
var MapFloat32s = io.MapFloat32s
var MapFloats = io.MapFloats
var MapInts = io.MapInts
var MarginalLogLikelihood = stats.MarginalLogLikelihood
var Matern32Kernel = stats.Matern32Kernel
var Matern52Kernel = stats.Matern52Kernel
var MatrixExp = linalg.MatrixExp
var MatrixLog = linalg.MatrixLog
var MatrixRank = linalg.MatrixRank
var MatrixSqrt = linalg.MatrixSqrt
var MaxPool1D = signal.MaxPool1D
var MaxPool2D = signal.MaxPool2D
var MaxPool3D = signal.MaxPool3D
var Mean = core.Mean
var Median = stats.Median
var MedianAbsoluteDeviation = stats.MedianAbsoluteDeviation
var MedianFilter = signal.MedianFilter
var MedianFilter2D = signal.MedianFilter2D
var Minimise = optim.Minimise
var MinimiseCMAES = optim.MinimiseCMAES
var MinimiseConstrained = optim.MinimiseConstrained
var MinimiseLBFGS = optim.MinimiseLBFGS
var MinimiseLinear = optim.MinimiseLinear
var MinimiseLinearRows = optim.MinimiseLinearRows
var MinimiseNonlinearConstrained = optim.MinimiseNonlinearConstrained
var MinimiseQP = optim.MinimiseQP
var MinimiseSimulatedAnnealing = optim.MinimiseSimulatedAnnealing
var MultivariateNormalDraws = stats.MultivariateNormalDraws
var MultivariateNormalLogDensity = stats.MultivariateNormalLogDensity
var NegativeBinomialCDF = stats.NegativeBinomialCDF
var NegativeBinomialPMF = stats.NegativeBinomialPMF
var NegativeBinomialQuantile = stats.NegativeBinomialQuantile
var NoncentralChiSquareCDF = stats.NoncentralChiSquareCDF
var NoncentralChiSquareDensity = stats.NoncentralChiSquareDensity
var NoncentralChiSquareQuantile = stats.NoncentralChiSquareQuantile
var NoncentralFCDF = stats.NoncentralFCDF
var NoncentralFQuantile = stats.NoncentralFQuantile
var NoncentralTCDF = stats.NoncentralTCDF
var NoncentralTQuantile = stats.NoncentralTQuantile
var NUFFTType1 = signal.NUFFTType1
var Nonzero = core.Nonzero
var NormalCDF = stats.NormalCDF
var NormalQuantile = stats.NormalQuantile
var NumWorkers = core.NumWorkers
var PCA = stats.PCA
var ParetoCDF = stats.ParetoCDF
var ParetoDensity = stats.ParetoDensity
var ParetoQuantile = stats.ParetoQuantile
var PeriodicKernel = stats.PeriodicKernel
var Pinverse = linalg.Pinverse
var PoissonCDF = stats.PoissonCDF
var PoissonDraws = stats.PoissonDraws
var PoissonQuantile = stats.PoissonQuantile
var PoissonRegression = stats.PoissonRegression
var PolynomialRoots = linalg.PolynomialRoots
var QR = linalg.QR
var Quantile = stats.Quantile
var QuantileRegression = stats.QuantileRegression
var RankFilter = signal.RankFilter
var RankFilter2D = signal.RankFilter2D
var RFFT = signal.RFFT
var SVD = linalg.SVD
var SVDComplex = linalg.SVDComplex
var SampleHMC = grad.SampleHMC
var SaveCSV = io.SaveCSV
var SaveCSVWriter = io.SaveCSVWriter
var SaveFITS = io.SaveFITS
var SaveFITSTable = io.SaveFITSTable
var SaveHDF5 = io.SaveHDF5
var SaveHDF5Text = io.SaveHDF5Text
var SaveNativeFloats = io.SaveNativeFloats
var SaveNetCDF = io.SaveNetCDF
var SavitzkyGolay = signal.SavitzkyGolay
var SelectARMA = signal.SelectARMA
var SchurComplex = linalg.SchurComplex
var SetNumCPU = core.SetNumCPU
var Solve = linalg.Solve
var SolveBoundaryCollocation = integrate.SolveBoundaryCollocation
var SolveCyclicTridiagonal = linalg.SolveCyclicTridiagonal
var SolvePoissonDirichlet = signal.SolvePoissonDirichlet
var SolvePoissonFEM3D = integrate.SolvePoissonFEM3D
var SolvePoissonNeumann = signal.SolvePoissonNeumann
var SolvePoissonPeriodic = signal.SolvePoissonPeriodic
var SolveRRQR = linalg.SolveRRQR
var SolveTikhonov = linalg.SolveTikhonov
var SolveTridiagonal = linalg.SolveTridiagonal
var SolveTruncated = linalg.SolveTruncated
var SpearmanRho = stats.SpearmanRho
var SpEigen = linalg.SpEigen
var SpEigenComplex = linalg.SpEigenComplex
var SpEigenGeneral = linalg.SpEigenGeneral
var SpEigenGeneralComplex = linalg.SpEigenGeneralComplex
var SpExpApply = linalg.SpExpApply
var SpLSMR = linalg.SpLSMR
var SpLSQR = linalg.SpLSQR
var SpSolve = linalg.SpSolve
var SpSolveBiCGSTAB = linalg.SpSolveBiCGSTAB
var SquaredExponentialKernel = stats.SquaredExponentialKernel
var Std = stats.Std
var SpSolveComplexCG = linalg.SpSolveComplexCG
var SpSolveComplexBiCGSTAB = linalg.SpSolveComplexBiCGSTAB
var StudentTCDF = stats.StudentTCDF
var StudentTDraws = stats.StudentTDraws
var StudentTQuantile = stats.StudentTQuantile
var SumKahan = signal.SumKahan
var TheilSenRegression = stats.TheilSenRegression
var Trace = core.Trace
var TraceComplex = core.TraceComplex
var Var = stats.Var
var VarSample = stats.VarSample
var WeibullCDF = stats.WeibullCDF
var WeibullDensity = stats.WeibullDensity
var WeibullQuantile = stats.WeibullQuantile
var WelchPSD = signal.WelchPSD
var WelchTTest = stats.WelchTTest
var WeightedLinearRegression = stats.WeightedLinearRegression
var WindowBartlett = signal.WindowBartlett
var WindowBlackman = signal.WindowBlackman
var WindowBlackmanHarris = signal.WindowBlackmanHarris
var WindowBox = signal.WindowBox
var WindowCosine = signal.WindowCosine
var WindowFlatTop = signal.WindowFlatTop
var WindowHamming = signal.WindowHamming
var WindowHann = signal.WindowHann
var WindowKaiser = signal.WindowKaiser
var EllipticK = core.EllipticK
var EllipticE = core.EllipticE
var EllipticPi = core.EllipticPi
var EllipticKScalar = core.EllipticKScalar
var EllipticFScalar = core.EllipticFScalar
var JacobiCDScalar = core.JacobiCDScalar
var JacobiSN = core.JacobiSN
var JacobiCN = core.JacobiCN
var JacobiDN = core.JacobiDN
var Hypergeometric2F1 = core.Hypergeometric2F1
var HaltonPoints = core.HaltonPoints
var SobolPoints = core.SobolPoints
var ButterworthBandPass = signal.ButterworthBandPass
var ButterworthBandStop = signal.ButterworthBandStop
var ButterworthHighPass = signal.ButterworthHighPass
var ButterworthLowPass = signal.ButterworthLowPass
var CauerBandPass = signal.CauerBandPass
var CauerBandStop = signal.CauerBandStop
var CauerHighPass = signal.CauerHighPass
var CauerLowPass = signal.CauerLowPass
var ChebyshevBandPass = signal.ChebyshevBandPass
var ChebyshevBandStop = signal.ChebyshevBandStop
var ChebyshevHighPass = signal.ChebyshevHighPass
var ChebyshevLowPass = signal.ChebyshevLowPass
var CSRFromCOO = linalg.CSRFromCOO
var CSCFromCOO = linalg.CSCFromCOO
var InverseChebyshevBandPass = signal.InverseChebyshevBandPass
var InverseChebyshevBandStop = signal.InverseChebyshevBandStop
var InverseChebyshevHighPass = signal.InverseChebyshevHighPass
var InverseChebyshevLowPass = signal.InverseChebyshevLowPass
var FilterApply = signal.FilterApply
var IntegrateODEEvents = integrate.IntegrateODEEvents
var LoadFITSTable = io.LoadFITSTable
var NewCubicSpline = linalg.NewCubicSpline
var NewSparseCholesky = linalg.NewSparseCholesky
var NewSparseLU = linalg.NewSparseLU
var NewSparseILU = linalg.NewSparseILU
var Pipe = linalg.Pipe
var IntegrateND = integrate.IntegrateND
var IntegrateHeat1D = integrate.IntegrateHeat1D
var IntegrateHeat2D = integrate.IntegrateHeat2D
var NewTriangleMesh2D = integrate.NewTriangleMesh2D
var GridTriangleMesh2D = integrate.GridTriangleMesh2D
var SolvePoissonFEM2D = integrate.SolvePoissonFEM2D
var IntegrateWave1D = integrate.IntegrateWave1D
var IntegrateWave2D = integrate.IntegrateWave2D
var Line = plot.Line
var LinearRegression = stats.LinearRegression
var Autocorrelate = signal.Autocorrelate
var CrossCorrelate = signal.CrossCorrelate
var PartialAutocorrelate = signal.PartialAutocorrelate
var DWT = signal.DWT
var IDWT = signal.IDWT
var CWT = signal.CWT
var MinimiseDifferentialEvolution = optim.MinimiseDifferentialEvolution
+24
View File
@@ -0,0 +1,24 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package tensor
import "testing"
func mustFromFloatsT(t *testing.T, v float64) *Array {
t.Helper()
a, err := FromFloats([]float64{v}, 1)
if err != nil {
t.Fatal(err)
}
return a
}
func FromFloatsMustT(t *testing.T, vals []float64) *Array {
t.Helper()
a, err := FromFloats(vals, len(vals))
if err != nil {
t.Fatal(err)
}
return a
}
+80
View File
@@ -0,0 +1,80 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package tensor
import (
"math"
"testing"
)
// The facade must forward every domain: one smoke call per package,
// through the re-exported names only.
func TestFacadeForwardsDomains(t *testing.T) {
// core: sum of a vector
a, _ := FromFloats([]float64{1, 2, 3}, 3)
if got := Sum(a).Int(); got != 6 {
t.Fatalf("Sum = %d", got)
}
// linalg: determinant
m, _ := FromFloats([]float64{1, 2, 3, 4}, 2, 2)
if det, _ := Det(m); det != -2 {
t.Fatalf("Det = %v", det)
}
// signal: DC of a constant via FFT
c, _ := FromFloats([]float64{2, 2, 2, 2}, 4)
spec, err := FFT(c)
if err != nil {
t.Fatal(err)
}
if real(spec.ComplexAt(0)) != 8 {
t.Fatalf("FFT DC = %v", spec.ComplexAt(0))
}
// integrate: BDF2 on decay
dy := func(t float64, y *Array) (*Array, error) {
return MulF(y, -1), nil
}
e, err := IntegrateBDF2(dy, 0, 1, mustFromFloatsT(t, 1), ODEOptions{MaxSteps: 100})
if err != nil {
t.Fatal(err)
}
if math.Abs(e.FloatAt(0)-1/math.E) > 1e-3 {
t.Fatalf("BDF2 decay = %v", e.FloatAt(0))
}
// stats: median
med, err := Median(FromFloatsMustT(t, []float64{3, 1, 2}))
if err != nil {
t.Fatal(err)
}
if med != 2 {
t.Fatalf("Median = %v", med)
}
// optim: Brent root of cos(x) - x
root, err := FindRoot(func(x float64) float64 { return math.Cos(x) - x }, 0, 1, 0)
if err != nil {
t.Fatal(err)
}
if math.Abs(root-0.7390851332151607) > 1e-9 {
t.Fatalf("root = %v", root)
}
// grad: simple backward
x := FromArray(mustFromFloatsT(t, 2), true)
loss, _ := x.Mul(x)
l, _ := loss.Sum()
if err := l.Backward(); err != nil {
t.Fatal(err)
}
if g := x.Grad().FloatAt(0); g != 4 {
t.Fatalf("d/dx x^2 at 2 = %v", g)
}
// core: the scalar elliptic functions
if got := EllipticKScalar(0.5); math.Abs(got-1.8540746773013719) > 1e-12 {
t.Fatalf("EllipticKScalar = %v", got)
}
if got := EllipticFScalar(0.3, 0.5); math.Abs(got-0.30225466857501754) > 1e-12 {
t.Fatalf("EllipticFScalar = %v", got)
}
if got := JacobiCDScalar(0.4, 0.5); math.Abs(got-0.9592196373527547) > 1e-12 {
t.Fatalf("JacobiCDScalar = %v", got)
}
}
+3
View File
@@ -0,0 +1,3 @@
module sourcedock.dev/petrbalvin/tensor
go 1.27.1
+313
View File
@@ -0,0 +1,313 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"slices"
ode "sourcedock.dev/petrbalvin/tensor/integrate"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Adjoint sensitivities of an initial value problem. Fitting the
// parameters of a differential equation to data asks for dL/dθ when
// the trajectory y(t; θ) is produced by an ODE solve, and
// differentiating through every solver step is neither necessary nor
// cheap. The adjoint method runs the dynamics once forward, then
// integrates the adjoint state λ(t) = ∂L/∂y(t) backward along
// λ' = −(∂f/∂y)ᵀλ from the loss gradient at the endpoint, carrying a
// per-parameter accumulator β' = −(∂f/∂θ)ᵀλ alongside; both
// Jacobian-vector products come from one automatic-differentiation
// backward pass per evaluation. The cost is one more ODE solve
// whatever the parameter count, which is what makes whole-trajectory
// fitting tractable.
// AdjointODE differentiates the solution of y' = f(t, y) at t1 with
// respect to the initial state and to the parameters the given
// function closes over. lossGrad is ∂L/∂y(t1), the seed the loss
// itself contributes; the return values are ∂L/∂y0 and, parallel to
// params, ∂L/∂θk, each shaped like its parameter. Backward-in-time
// problems (t1 < t0) work.
//
// The forward trajectory is recorded at the adaptive solver's accepted
// steps and handed to the backward pass through cubic Hermite
// interpolation, fourth-order accurate like the Dormand-Prince pair
// that produced it; the augmented adjoint system is integrated by the
// same adaptive solver in reverse. The parameters' own accumulated
// gradients are left untouched: every evaluation runs a reverse pass
// that computes into a local map and commits nothing, and the answers
// travel out as return values instead.
//
// A nil function, a non-vector start or loss seed, a parameter that
// does not require grad, and an f that ignores its state or returns a
// wrong shape are errors, never silent zeros.
func AdjointODE(f func(t float64, y *Tensor) (*Tensor, error),
params []*Tensor, t0, t1 float64, y0 *core.Array,
lossGrad *core.Array, opts ode.ODEOptions) (*core.Array, []*core.Array, error) {
const name = "AdjointODE"
if f == nil {
return nil, nil, errf("%s: f must not be nil", name)
}
if y0 == nil || y0.NDim() != 1 || y0.Len() == 0 {
return nil, nil, errf("%s: the state must be a non-empty vector", name)
}
if y0.Dtype() == core.Complex {
return nil, nil, errf("%s: complex states are not supported", name)
}
dim := y0.Len()
if lossGrad == nil || lossGrad.NDim() != 1 || lossGrad.Len() != dim {
return nil, nil, errf("%s: lossGrad must be a vector of length %d", name, dim)
}
sizes := make([]int, len(params))
total := 0
for k, p := range params {
if p == nil || !p.RequiresGrad() {
return nil, nil, errf("%s: parameter %d does not require grad", name, k)
}
if p.Data().NDim() == 0 || p.Data().Len() == 0 {
return nil, nil, errf("%s: parameter %d must be a non-empty tensor of rank at least 1", name, k)
}
sizes[k] = p.Data().Len()
total += sizes[k]
}
if t0 == t1 {
// No dynamics: the endpoint is the initial state.
seed := flatFloats(lossGrad)
gradY0, err := core.FromFloats(seed, dim)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
return gradY0, zeroBlocks(params, sizes), nil
}
// Forward pass with the dynamics evaluated for value only.
forward := func(t float64, ya *core.Array) (*core.Array, error) {
out, err := f(t, FromArray(ya, false))
if err != nil {
return nil, errf("%s: %w", name, err)
}
return out.Data(), nil
}
times, states, err := ode.IntegrateODESteps(forward, t0, t1, y0, opts)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
// Cubic Hermite interpolation wants the derivative of the dynamics
// at every recorded node; one detached pass supplies them all.
slopes := make([][]float64, len(times))
for i, node := range states {
out, err := f(times[i], FromArray(node, false))
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
if out.Data().NDim() != 1 || out.Data().Len() != dim {
return nil, nil, errf("%s: f returned shape %s, want a vector of length %d",
name, prettyShape(out.Data().Shape()), dim)
}
slopes[i] = flatFloats(out.Data())
}
trace := newODETrace(times, states, slopes, dim)
vjp := func(t float64, y []float64, lam []float64) ([]float64, []float64, error) {
data, err := core.FromFloats(y, dim)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
yLeaf := FromArray(data, true)
out, err := f(t, yLeaf)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
if out.Data().NDim() != 1 || out.Data().Len() != dim {
return nil, nil, errf("%s: f returned shape %s, want a vector of length %d",
name, prettyShape(out.Data().Shape()), dim)
}
seedData, err := core.FromFloats(lam, dim)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
weighted, err := out.Mul(FromArray(seedData, false))
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
summed, err := weighted.Sum()
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
// The Jacobian-vector products come from the reverse pass
// alone: reverseGrads computes every reached tensor's gradient
// into a local map and commits nothing, so no leaf the dynamics
// close over, in params or out, is written at all, and the
// callers' own accumulated gradients survive untouched.
grads, err := summed.reverseGrads()
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
g := grads[yLeaf]
if g == nil || g.Len() != dim {
return nil, nil, errf("%s: f did not yield a state gradient of length %d", name, dim)
}
gY := flatFloats(g)
gTheta := make([]float64, 0, total)
for k, p := range params {
now := grads[p]
if now == nil {
return nil, nil, errf("%s: parameter %d is disconnected from the dynamics, no gradient flowed", name, k)
}
nowFloats := flatFloats(now)
gTheta = append(gTheta, nowFloats...)
}
return gY, gTheta, nil
}
rhs := func(t float64, z *core.Array) (*core.Array, error) {
lam := make([]float64, dim)
if !z.Strided() && z.Dtype() == core.Float {
copy(lam, z.RawFloats()[:dim])
} else {
for i := range lam {
lam[i] = z.FloatAt(i)
}
}
yArr, err := trace.at(t)
if err != nil {
return nil, err
}
gY, gTheta, err := vjp(t, yArr, lam)
if err != nil {
return nil, err
}
out := make([]float64, dim+total)
for i := range gY {
out[i] = -gY[i]
}
for i := range gTheta {
out[dim+i] = -gTheta[i]
}
return core.FromFloats(out, dim+total)
}
zStart := make([]float64, dim+total)
if !lossGrad.Strided() && lossGrad.Dtype() == core.Float {
copy(zStart, lossGrad.RawFloats()[:dim])
} else {
for i := range dim {
zStart[i] = lossGrad.FloatAt(i)
}
}
zEnd, err := ode.IntegrateODE(rhs, t1, t0, mustFromFloats(zStart, dim+total), opts)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
seed := make([]float64, dim)
for i := range seed {
seed[i] = zEnd.FloatAt(i)
}
gradY0, err := core.FromFloats(seed, dim)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
blocks := make([]*core.Array, len(params))
offset := dim
for k, p := range params {
block := make([]float64, sizes[k])
if !zEnd.Strided() && zEnd.Dtype() == core.Float {
copy(block, zEnd.RawFloats()[offset:offset+sizes[k]])
} else {
for j := range block {
block[j] = zEnd.FloatAt(offset + j)
}
}
shaped, err := core.FromFloats(block, p.Data().Shape()...)
if err != nil {
return nil, nil, errf("%s: %w", name, err)
}
blocks[k] = shaped
offset += sizes[k]
}
return gradY0, blocks, nil
}
// zeroBlocks builds zero arrays shaped like each parameter, the
// degenerate answer of a span with no dynamics.
func zeroBlocks(params []*Tensor, sizes []int) []*core.Array {
blocks := make([]*core.Array, len(params))
for k, p := range params {
zeros := make([]float64, sizes[k])
blocks[k], _ = core.FromFloats(zeros, p.Data().Shape()...)
}
return blocks
}
// mustFromFloats wraps a plain construction that cannot fail: the
// length always matches the single-dimension shape.
func mustFromFloats(vals []float64, n int) *core.Array {
a, _ := core.FromFloats(vals, n)
return a
}
// odeTrace is a recorded forward trajectory, kept ascending in time,
// with the dynamics' derivative at every node so the interpolation is
// cubic Hermite.
type odeTrace struct {
times []float64
states [][]float64
slopes [][]float64
dim int
// buf is the interpolation buffer at reuses: the solver calls at
// once per right-hand-side evaluation, so the array is borrowed
// rather than allocated each time. Callers must consume the answer
// before the next call.
buf []float64
}
// newODETrace flattens and time-orders the recorded nodes.
func newODETrace(times []float64, states []*core.Array, slopes [][]float64, dim int) *odeTrace {
tr := &odeTrace{times: append([]float64{}, times...), slopes: slopes, dim: dim}
tr.states = make([][]float64, len(states))
for i, s := range states {
tr.states[i] = flatFloats(s)
}
if len(times) > 1 && times[1] < times[0] {
// A backward-in-time pass records descending nodes; flip to
// ascending so the search below has one convention.
for i, j := 0, len(tr.times)-1; i < j; i, j = i+1, j-1 {
tr.times[i], tr.times[j] = tr.times[j], tr.times[i]
tr.states[i], tr.states[j] = tr.states[j], tr.states[i]
tr.slopes[i], tr.slopes[j] = tr.slopes[j], tr.slopes[i]
}
}
return tr
}
// at interpolates the state at t by cubic Hermite over the enclosing
// recorded interval, fourth-order accurate in the step size and never
// leaving the recorded span. A trace with fewer than two nodes offers
// no interval to interpolate over, so the accessor refuses instead of
// indexing out of range. The returned slice belongs to the trace and
// stays valid only until the next call.
func (tr *odeTrace) at(t float64) ([]float64, error) {
n := len(tr.times)
if n < 2 {
return nil, errf("AdjointODE: the solver recorded %d trajectory nodes, the adjoint interpolation needs at least two", n)
}
idx, _ := slices.BinarySearch(tr.times, t)
i := min(max(idx-1, 0), n-2)
h := tr.times[i+1] - tr.times[i]
s := (t - tr.times[i]) / h
h00 := s*s*(2*s-3) + 1
h10 := s * (s - 1) * (s - 1)
h01 := s * s * (3 - 2*s)
h11 := s * s * (s - 1)
if len(tr.buf) != tr.dim {
tr.buf = make([]float64, tr.dim)
}
out := tr.buf
si, si1 := tr.states[i], tr.states[i+1]
li, li1 := tr.slopes[i], tr.slopes[i+1]
for j := range out {
out[j] = h00*si[j] + h01*si1[j] +
h*(h10*li[j]+h11*li1[j])
}
return out, nil
}
+349
View File
@@ -0,0 +1,349 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
ode "sourcedock.dev/petrbalvin/tensor/integrate"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// decayWith builds f for y' = −θy with θ a scalar parameter leaf.
func decayWith(theta *Tensor) func(t float64, y *Tensor) (*Tensor, error) {
return func(t float64, y *Tensor) (*Tensor, error) {
rate, err := theta.Neg()
if err != nil {
return nil, err
}
return y.Mul(rate)
}
}
// scalarVector returns a length-1 real array.
func scalarVector(t *testing.T, v float64) *core.Array {
t.Helper()
a, err := core.FromFloats([]float64{v}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
return a
}
// TestAdjointODEDecay differentiates y' = −θy with L = y(1): the exact
// sensitivities are dL/dy0 = e^{−θ} and dL/dθ = −e^{−θ}, and a run
// with no parameters at all must still answer the initial-state one.
func TestAdjointODEDecay(t *testing.T) {
const want = 0.4965853037914095 // e^{−0.7}
theta, err := FromFloat64s([]float64{0.7}, true, 1)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1,
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 {
t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want)
}
if len(blocks) != 1 || math.Abs(blocks[0].FloatAt(0)+want) > 1e-6 {
t.Fatalf("dL/dθ = %v, want −%.14g", blocks[0].FloatAt(0), want)
}
// The parameter's own accumulated gradient must be untouched.
if theta.Grad() != nil {
t.Fatal("AdjointODE must leave the parameters' gradients untouched")
}
solo, soloBlocks, err := AdjointODE(decayWith(theta), nil, 0, 1,
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE without parameters: %v", err)
}
if len(soloBlocks) != 0 || math.Abs(solo.FloatAt(0)-want) > 1e-6 {
t.Fatalf("parameter-free run: dL/dy0 = %.14g with %d blocks", solo.FloatAt(0), len(soloBlocks))
}
}
// TestAdjointODEOscillator differentiates the harmonic oscillator
// y” = −ω²y with L = y(1) and y0 = (1, 0): dL/dω = −sin(ω),
// dL/dy0 = (cos ω, sin ω/ω).
func TestAdjointODEOscillator(t *testing.T) {
omega, err := FromFloat64s([]float64{1.3}, true, 1)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(t float64, y *Tensor) (*Tensor, error) {
u, err := y.Slice(0, 0, 1)
if err != nil {
return nil, err
}
v, err := y.Slice(0, 1, 2)
if err != nil {
return nil, err
}
sq, err := omega.Mul(omega)
if err != nil {
return nil, err
}
acc, err := u.Mul(sq)
if err != nil {
return nil, err
}
drag, err := acc.Scale(-1)
if err != nil {
return nil, err
}
return v.Concat(drag, 0)
}
y0, err := core.FromFloats([]float64{1, 0}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
seed, err := core.FromFloats([]float64{1, 0}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
gradY0, blocks, err := AdjointODE(f, []*Tensor{omega}, 0, 1, y0, seed,
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
if math.Abs(gradY0.FloatAt(0)-math.Cos(1.3)) > 1e-6 {
t.Fatalf("dL/du0 = %.14g, want cos(1.3) = %.14g", gradY0.FloatAt(0), math.Cos(1.3))
}
if math.Abs(gradY0.FloatAt(1)-math.Sin(1.3)/1.3) > 1e-6 {
t.Fatalf("dL/dv0 = %.14g, want sin(1.3)/1.3 = %.14g",
gradY0.FloatAt(1), math.Sin(1.3)/1.3)
}
if math.Abs(blocks[0].FloatAt(0)+math.Sin(1.3)) > 1e-6 {
t.Fatalf("dL/dω = %.14g, want −sin(1.3) = %.14g", blocks[0].FloatAt(0), -math.Sin(1.3))
}
}
// TestAdjointODEBackwardTime runs the forward pass itself backwards
// (t1 < t0): y(t) = e^{−θ(t−1)} from y(1) = 1 has y(0) = e^θ and both
// sensitivities equal e^θ.
func TestAdjointODEBackwardTime(t *testing.T) {
const want = 1.6487212707001282 // e^{0.5}
theta, _ := FromFloat64s([]float64{0.5}, true, 1)
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 1, 0,
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 {
t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want)
}
if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 {
t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want)
}
}
// TestAdjointODEFiniteDifference checks a nonlinear two-parameter
// system against central differences of the forward solve itself,
// perturbing the parameter leaves around the adjoint run.
func TestAdjointODEFiniteDifference(t *testing.T) {
th1, _ := FromFloat64s([]float64{0.8}, true, 1)
th2, _ := FromFloat64s([]float64{1.1}, true, 1)
f := func(t float64, y *Tensor) (*Tensor, error) {
y1, err := y.Slice(0, 0, 1)
if err != nil {
return nil, err
}
y2, err := y.Slice(0, 1, 2)
if err != nil {
return nil, err
}
drag, err := y1.Mul(th1)
if err != nil {
return nil, err
}
drive, err := y2.Mul(th2)
if err != nil {
return nil, err
}
r1, err := drive.Sub(drag)
if err != nil {
return nil, err
}
pool, err := y1.Mul(y2)
if err != nil {
return nil, err
}
r2, err := pool.Scale(-1)
if err != nil {
return nil, err
}
return r1.Concat(r2, 0)
}
y0, _ := core.FromFloats([]float64{1, 0.5}, 2)
seed, _ := core.FromFloats([]float64{1, 2}, 2)
gradY0, blocks, err := AdjointODE(f, []*Tensor{th1, th2}, 0, 1, y0, seed,
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-13})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
// Central differences on the loss y1(1) + 2·y2(1), with the leaf
// data swapped out for the perturbed values.
forwardLoss := func() float64 {
end, err := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) {
out, err := f(t, FromArray(ya, false))
if err != nil {
return nil, err
}
return out.Data(), nil
}, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15})
if err != nil {
t.Fatalf("forward solve: %v", err)
}
return end.FloatAt(0) + 2*end.FloatAt(1)
}
perturb := func(p *Tensor, eps float64) {
swapped, _ := core.FromFloats([]float64{p.Data().FloatAt(0) + eps}, 1)
p.ReplaceWith(swapped)
}
restore := func(p *Tensor, v float64) {
orig, _ := core.FromFloats([]float64{v}, 1)
p.ReplaceWith(orig)
}
const eps = 1e-5
for k, p := range []*Tensor{th1, th2} {
orig := p.Data().FloatAt(0)
perturb(p, eps)
up := forwardLoss()
perturb(p, -2*eps)
down := forwardLoss()
restore(p, orig)
fd := (up - down) / (2 * eps)
got := blocks[k].FloatAt(0)
if math.Abs(got-fd) > 1e-3*math.Max(1, math.Abs(fd)) {
t.Fatalf("dL/dθ%d: adjoint %.8g, finite difference %.8g", k+1, got, fd)
}
}
if gradY0.Len() != 2 {
t.Fatalf("dL/dy0 has length %d, want 2", gradY0.Len())
}
}
// TestAdjointODEDegenerateSpan returns the loss seed unchanged and
// zero parameter gradients when the span carries no dynamics.
func TestAdjointODEDegenerateSpan(t *testing.T) {
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0.5, 0.5,
scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
if gradY0.FloatAt(0) != 2 {
t.Fatalf("dL/dy0 = %.14g, want the seed 2", gradY0.FloatAt(0))
}
if blocks[0].FloatAt(0) != 0 {
t.Fatalf("dL/dθ = %.14g, want 0", blocks[0].FloatAt(0))
}
}
// TestAdjointODEErrors pins the validation contract.
func TestAdjointODEErrors(t *testing.T) {
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
y0 := scalarVector(t, 1)
seed := scalarVector(t, 1)
rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2)
if _, _, err := AdjointODE(nil, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a nil function")
}
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, nil, seed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a nil state")
}
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, rank2, seed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, nil, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a nil loss seed")
}
badSeed, _ := core.FromFloats([]float64{1, 1}, 2)
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, badSeed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a loss seed of the wrong length")
}
frozen := FromArray(scalarVector(t, 0.7), false)
if _, _, err := AdjointODE(decayWith(frozen), []*Tensor{frozen}, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a parameter that does not require grad")
}
wrongShape := func(t float64, y *Tensor) (*Tensor, error) {
three, _ := core.FromFloats([]float64{1, 2, 3}, 3)
return FromArray(three, false), nil
}
if _, _, err := AdjointODE(wrongShape, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error for a wrong-shaped derivative")
}
ignoresState := func(t float64, y *Tensor) (*Tensor, error) {
one, _ := core.FromFloats([]float64{1}, 1)
return FromArray(one, false), nil
}
if _, _, err := AdjointODE(ignoresState, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
t.Fatal("expected an error when f ignores its state")
}
}
// TestAdjointODERestoresGradientsOnError pins the error-path contract:
// when f consumes the parameter but ignores the state, no state
// gradient can flow and the run errors, and the caller's own
// accumulated parameter gradient must come back untouched instead of
// being polluted by the aborted pass: the reverse sweep is not run at
// all, so nothing writes the leaf gradients on the way out.
func TestAdjointODERestoresGradientsOnError(t *testing.T) {
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
preset, _ := core.FromFloats([]float64{3}, 1)
theta.SetGrad(preset)
f := func(t float64, y *Tensor) (*Tensor, error) {
return theta.Scale(2)
}
if _, _, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1),
scalarVector(t, 1), ode.ODEOptions{}); err == nil {
t.Fatal("expected an error when f ignores its state")
}
if theta.Grad() == nil || theta.Grad().FloatAt(0) != 3 {
t.Fatalf("the parameter's gradient was not restored: %v, want 3", theta.Grad())
}
}
// TestAdjointODEDisconnectedParameter pins that a parameter the
// dynamics never touch is an error, not a silent zero gradient block.
func TestAdjointODEDisconnectedParameter(t *testing.T) {
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
f := func(t float64, y *Tensor) (*Tensor, error) {
return y.Scale(2)
}
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1),
scalarVector(t, 1), ode.ODEOptions{})
if err == nil {
t.Fatal("expected an error for a parameter disconnected from the dynamics")
}
if blocks != nil {
t.Fatalf("an errored run returned blocks: %v", blocks)
}
}
// TestAdjointODESuccessKeepsPresetGradient pins the same restore on
// the success path: a preset accumulated gradient survives the run.
func TestAdjointODESuccessKeepsPresetGradient(t *testing.T) {
theta, _ := FromFloat64s([]float64{0.5}, true, 1)
preset, _ := core.FromFloats([]float64{7}, 1)
theta.SetGrad(preset)
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1,
scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
const want = 0.6065306597126334 // e^{−0.5}
if math.Abs(gradY0.FloatAt(0)-2*want) > 1e-6 {
t.Fatalf("dL/dy0 = %.14g, want 2e^{−0.5}", gradY0.FloatAt(0))
}
if theta.Grad() == nil || theta.Grad().FloatAt(0) != 7 {
t.Fatalf("the preset gradient did not survive a successful run: %v", theta.Grad())
}
if blocks[0].FloatAt(0) >= 0 {
t.Fatalf("dL/dθ = %v, want negative", blocks[0].FloatAt(0))
}
}
+918
View File
@@ -0,0 +1,918 @@
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Backward-coverage pins for the graph kernels: every test here checks
// a committed leaf gradient against a closed-form identity or a central
// difference of the summed loss, on combinations the op-level tests do
// not build.
// totalLoss sums a non-scalar loss so a central difference matches the
// all-ones seed Backward uses.
func totalLoss(l *Tensor) float64 {
s := 0.0
d := l.Data()
for i := range d.Len() {
s += d.FloatAt(i)
}
return s
}
func abs2c(z complex128) float64 { return real(z)*real(z) + imag(z)*imag(z) }
func mustRecover(vals []float64, sh ...int) *core.Array {
a, _ := core.FromFloats(vals, sh...)
return a
}
func mustRecoverComplex(zs []complex128, shape ...int) *core.Array {
if len(shape) == 0 {
shape = []int{len(zs)}
}
w := append([]complex128(nil), zs...)
a, _ := core.ComplexFromArray(w, shape...)
return a
}
// probeCheckCentral builds the loss from base, backpropagates once and
// compares every element of the committed gradient against central
// differences of the summed loss.
func probeCheckCentral(t *testing.T, base *Tensor, name string, build func() (*Tensor, error), tol float64) {
t.Helper()
base.ZeroGrad()
loss, err := build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
grads := flatFloats(base.Grad())
orig := flatFloats(base.Data())
sh := base.Data().Shape()
for i := range orig {
const h = 1e-6
plus := append([]float64(nil), orig...)
plus[i] += h
base.ReplaceWith(mustRecover(plus, sh...))
lp, err := build()
if err != nil {
t.Fatal(err)
}
minus := append([]float64(nil), orig...)
minus[i] -= h
base.ReplaceWith(mustRecover(minus, sh...))
lm, err := build()
if err != nil {
t.Fatal(err)
}
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
if math.Abs(grads[i]-num) > tol*math.Max(1, math.Abs(num)) {
t.Errorf("%s grad[%d] = %g, central %g", name, i, grads[i], num)
}
base.ReplaceWith(mustRecover(orig, sh...))
}
}
// TestFFTParsevalGradient pins the FFT backward through Parseval's
// theorem: L = sum |FFT(x)|^2 has gradient n·2x for real x, because the
// unnormalised forward DFT scales the energy by n.
func TestFFTParsevalGradient(t *testing.T) {
const n = 16
vals := make([]float64, n)
for i := range vals {
vals[i] = math.Sin(0.7*float64(i)) + 0.3*float64(i%4)
}
x, _ := FromFloat64s(vals, true, n)
f, _ := x.FFT()
abs2, _ := f.Abs2()
loss, _ := abs2.Sum()
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
g := x.Grad()
for i := range n {
want := 2 * float64(n) * vals[i]
if got := g.FloatAt(i); math.Abs(got-want) > 1e-8*math.Max(1, math.Abs(want)) {
t.Fatalf("Parseval: g[%d] = %g, want %g", i, got, want)
}
}
}
// TestFFT2ParsevalGradient is Parseval's identity for the rank-2
// transform, where the backward's scale is the full element count.
func TestFFT2ParsevalGradient(t *testing.T) {
vals := []float64{1, 2, -1, 0.5, 3, -2, 0.25, 1.5, -0.5, 2, 1, -3}
x, _ := FromFloat64s(vals, true, 3, 4)
f, err := x.FFT2()
if err != nil {
t.Fatal(err)
}
a, err := f.Abs2()
if err != nil {
t.Fatal(err)
}
loss, err := a.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
for i := range vals {
want := 2 * float64(12) * vals[i]
if got := x.Grad().FloatAt(i); math.Abs(got-want) > 1e-8*math.Max(1, math.Abs(want)) {
t.Fatalf("FFT2 Parseval g[%d] = %g, want %g", i, got, want)
}
}
}
// TestSpectralHalfSpectrumBackward pins the RFFT and IRFFT backwards
// against central differences, the half-spectrum combinatorics
// included: the doubled real part on the mirrored bins and the halved
// self-mirrored bins of the inverse.
func TestSpectralHalfSpectrumBackward(t *testing.T) {
x, _ := FromFloat64s([]float64{1, -2, 0.5, 3, -1.5, 0.25, 2, -0.75}, true, 8)
buildF := func() (*Tensor, error) {
h, err := x.RFFT()
if err != nil {
return nil, err
}
a, err := h.Abs2()
if err != nil {
return nil, err
}
return a.Sum()
}
probeCheckCentral(t, x, "rfft", buildF, 1e-4)
// IRFFT: the leaf is the half spectrum; the loss is the summed
// square of the real signal. The complex leaf's gradient is checked
// against a Wirtinger central difference of the real loss:
// dL/dzbar_j = (dL/dRe_j + i dL/dIm_j)/2, each part probed by
// perturbing that component.
zr := []float64{3, -1, 0.5, 2, 1.5}
zi := []float64{0, 1, -0.5, 0.25, -1}
zs := make([]complex128, 5)
for i := range zs {
zs[i] = complex(zr[i], zi[i])
}
za, _ := core.ComplexFromArray(zs, 5)
z := FromArray(za, true)
buildI := func() (*Tensor, error) {
sig, err := z.IRFFT(8)
if err != nil {
return nil, err
}
sq, err := sig.Pow(2)
if err != nil {
return nil, err
}
return sq.Sum()
}
z.ZeroGrad()
loss, err := buildI()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
const h = 1e-6
for j := range 5 {
perturb := func(dr, di float64) float64 {
ws := append([]complex128(nil), zs...)
ws[j] = complex(zr[j]+dr, zi[j]+di)
z.ReplaceWith(mustRecoverComplex(ws, 5))
l, err := buildI()
if err != nil {
t.Fatal(err)
}
return totalLoss(l)
}
dr := (perturb(h, 0) - perturb(-h, 0)) / (2 * h)
di := (perturb(0, h) - perturb(0, -h)) / (2 * h)
want := complex(0.5*dr, 0.5*di)
if got := z.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
t.Errorf("irfft grad[%d] = %g, want %g", j, got, want)
}
z.ReplaceWith(mustRecoverComplex(zs, 5))
}
}
// TestCompositeGraphBackward walks a graph spanning slice, matmul,
// tanh, broadcast, abs2 and an axis reduction, and checks both leaves'
// gradients against central differences of the summed loss.
func TestCompositeGraphBackward(t *testing.T) {
a, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
w, _ := FromFloat64s([]float64{0.5, -0.25, 0.125, 1, -2, 0.75}, true, 2, 3)
build := func() (*Tensor, error) {
s, err := a.Slice(1, 1, 3) // (2,2)
if err != nil {
return nil, err
}
mm, err := s.MatMul(w)
if err != nil {
return nil, err
}
th, err := mm.Tanh()
if err != nil {
return nil, err
}
bc, err := th.BroadcastTo(2, 2, 3)
if err != nil {
return nil, err
}
sq, err := bc.Abs2()
if err != nil {
return nil, err
}
return sq.SumAxis(2)
}
numCheck := func(base *Tensor, name string, grads []float64) {
t.Helper()
orig := flatFloats(base.Data())
sh := base.Data().Shape()
for i := range orig {
const h = 1e-6
plus := append([]float64(nil), orig...)
plus[i] += h
base.ReplaceWith(mustRecover(plus, sh...))
lp, err := build()
if err != nil {
t.Fatal(err)
}
minus := append([]float64(nil), orig...)
minus[i] -= h
base.ReplaceWith(mustRecover(minus, sh...))
lm, err := build()
if err != nil {
t.Fatal(err)
}
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
t.Fatalf("%s[%d]: backward %g, central difference %g", name, i, grads[i], num)
}
base.ReplaceWith(mustRecover(orig, sh...))
}
}
// Backward accumulates, so every read starts from a clean slate.
a.ZeroGrad()
w.ZeroGrad()
loss, err := build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
numCheck(a, "a", flatFloats(a.Grad()))
a.ZeroGrad()
w.ZeroGrad()
loss, err = build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
numCheck(w, "w", flatFloats(w.Grad()))
}
// TestMatMulBatchedBackward checks the stacked product rule of
// MatMulBatched against central differences on both operands.
func TestMatMulBatchedBackward(t *testing.T) {
a, _ := FromFloat64s([]float64{1, 2, 3, 4, 5, 6, 7, 8}, true, 2, 2, 2)
b, _ := FromFloat64s([]float64{0.5, -1, 2, 0.25, 1.5, -0.5, -2, 1}, true, 2, 2, 2)
build := func() (*Tensor, error) {
c, err := a.MatMulBatched(b)
if err != nil {
return nil, err
}
sq, err := c.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
a.ZeroGrad()
b.ZeroGrad()
loss, err := build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
tensor *Tensor
name string
}{{a, "a"}, {b, "b"}} {
grads := flatFloats(tc.tensor.Grad())
orig := flatFloats(tc.tensor.Data())
sh := tc.tensor.Data().Shape()
for i := range orig {
const h = 1e-6
plus := append([]float64(nil), orig...)
plus[i] += h
tc.tensor.ReplaceWith(mustRecover(plus, sh...))
lp, err := build()
if err != nil {
t.Fatal(err)
}
minus := append([]float64(nil), orig...)
minus[i] -= h
tc.tensor.ReplaceWith(mustRecover(minus, sh...))
lm, err := build()
if err != nil {
t.Fatal(err)
}
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
if math.Abs(grads[i]-num) > 1e-4*math.Max(1, math.Abs(num)) {
t.Fatalf("%s[%d]: backward %g, central difference %g", tc.name, i, grads[i], num)
}
tc.tensor.ReplaceWith(mustRecover(orig, sh...))
}
}
}
// TestStridedOperandGradientMatchesDenseTwin pins the graph's stride
// discipline: a MatMul over a sliced view and the same product over a
// dense twin must commit identical gradients. The view's leaf receives
// the dense twin's gradient scattered into the sliced span, zeros
// elsewhere.
func TestStridedOperandGradientMatchesDenseTwin(t *testing.T) {
wVals := []float64{0.5, -0.25, 0.125, 1, -2, 0.75}
// The strided run: leaf (2,4) -> slice -> matmul.
full, _ := FromFloat64s([]float64{9, 1, -2, 3, 9, 5, -6, 9}, true, 2, 4)
s, err := full.Slice(1, 1, 4) // (2,3) strided view holding 1,-2,3 / 5,-6,9
if err != nil {
t.Fatal(err)
}
w, _ := FromFloat64s(wVals, true, 3, 2)
mm, err := s.MatMul(w)
if err != nil {
t.Fatal(err)
}
sq, err := mm.Abs2()
if err != nil {
t.Fatal(err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
if full.Grad() == nil || w.Grad() == nil {
t.Fatal("both leaves must receive a gradient")
}
gs := flatFloats(full.Grad())
wStrided := flatFloats(w.Grad())
full.ZeroGrad()
w.ZeroGrad()
// The dense twin.
a2, _ := FromFloat64s(flatFloats(s.Data()), true, 2, 3)
w2, _ := FromFloat64s(wVals, true, 3, 2)
mm2, err := a2.MatMul(w2)
if err != nil {
t.Fatal(err)
}
sq2, err := mm2.Abs2()
if err != nil {
t.Fatal(err)
}
loss2, err := sq2.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss2.Backward(); err != nil {
t.Fatal(err)
}
aDense := flatFloats(a2.Grad())
wDense := flatFloats(w2.Grad())
for i := range wStrided {
if wStrided[i] != wDense[i] {
t.Fatalf("w gradient %v over the strided operand, want the dense twin's %v", wStrided, wDense)
}
}
if gs[0] != 0 || gs[4] != 0 {
t.Fatalf("columns outside the slice carry %g and %g, want 0", gs[0], gs[4])
}
if gs[1] != aDense[0] || gs[2] != aDense[1] || gs[3] != aDense[2] ||
gs[5] != aDense[3] || gs[6] != aDense[4] || gs[7] != aDense[5] {
t.Fatalf("leaf gradient %v does not scatter the dense gradient %v into the slice", gs, aDense)
}
}
// TestHessianMatchesAnalyticForm differentiates
// f(x, y) = x^3 y + exp(x) log(y+2) twice and compares the answer with
// the closed form:
//
// dxx = 6xy + exp(x)log(y+2), dxy = 3x^2 + exp(x)/(y+2),
// dyy = -exp(x)/(y+2)^2.
func TestHessianMatchesAnalyticForm(t *testing.T) {
f := func(v *Tensor) (*Tensor, error) {
x, err := v.Slice(0, 0, 1)
if err != nil {
return nil, err
}
y, err := v.Slice(0, 1, 2)
if err != nil {
return nil, err
}
x3, err := x.Pow(3)
if err != nil {
return nil, err
}
t1, err := x3.Mul(y)
if err != nil {
return nil, err
}
e, err := x.Exp()
if err != nil {
return nil, err
}
y2, err := y.Add(FromArray(mustRecover([]float64{2}, 1), false))
if err != nil {
return nil, err
}
lg, err := y2.Log()
if err != nil {
return nil, err
}
t2, err := e.Mul(lg)
if err != nil {
return nil, err
}
return t1.Add(t2)
}
x, _ := FromFloat64s([]float64{0.4, 1.7}, true, 2)
h, err := Hessian(f, x, HessianOptions{Step: 1e-5})
if err != nil {
t.Fatal(err)
}
xx, yy := 0.4, 1.7
ex, ly := math.Exp(xx), math.Log(yy+2)
want := [4]float64{
6*xx*yy + ex*ly, 3*xx*xx + ex/(yy+2),
3*xx*xx + ex/(yy+2), -ex / ((yy + 2) * (yy + 2)),
}
for i := range 4 {
if math.Abs(h.FloatAt(i)-want[i]) > 1e-4*math.Max(1, math.Abs(want[i])) {
t.Fatalf("Hessian[%d] = %g, want %g", i, h.FloatAt(i), want[i])
}
}
}
// TestTripleUseLeafGradient folds a leaf through three nodes and pins
// the accumulated gradient: dL/dz = 2(v^2+v)(2v+1) for
// L = (v^2+v)^2.
func TestTripleUseLeafGradient(t *testing.T) {
for _, v := range []float64{0.5, -1.25, 2} {
z, _ := FromFloat64s([]float64{v}, true, 1)
z2, _ := z.Mul(z)
s, err := z2.Add(z)
if err != nil {
t.Fatal(err)
}
sq, err := s.Pow(2)
if err != nil {
t.Fatal(err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
want := 2 * (v*v + v) * (2*v + 1)
if got := z.Grad().FloatAt(0); math.Abs(got-want) > 1e-12*math.Max(1, math.Abs(want)) {
t.Fatalf("triple use at z=%g: g = %g, want %g", v, got, want)
}
}
}
// TestComplexAbs2LeafGradient pins the Wirtinger gradient of a summed
// |z|^2 loss: the leaf receives z itself.
func TestComplexAbs2LeafGradient(t *testing.T) {
zr := []float64{0.3, -1.2, 0.7}
zi := []float64{-0.4, 0.8, 1.1}
zs := make([]complex128, 3)
for i := range zs {
zs[i] = complex(zr[i], zi[i])
}
za, _ := core.ComplexFromArray(zs, 3)
z := FromArray(za, true)
abs2, err := z.Abs2()
if err != nil {
t.Fatal(err)
}
loss, err := abs2.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
g := z.Grad()
for i := range 3 {
want := complex(zr[i], zi[i])
if got := g.ComplexAt(i); got != want {
t.Fatalf("|z|^2 leaf grad[%d] = %g, want %g", i, got, want)
}
}
}
// TestComplexMatMulBackward checks the Wirtinger adjoint of MatMul
// against component-wise central differences of a real loss. The leaf
// gradient is dL/dzbar = (dL/dRe + i dL/dIm)/2 under the package
// convention.
func TestComplexMatMulBackward(t *testing.T) {
az := []complex128{complex(1, 0.5), complex(-0.5, 2), complex(0.25, -1), complex(2, 0.25)}
bz := []complex128{complex(0.5, 1), complex(-1, 0.25), complex(1.5, -0.5), complex(0.75, 1.25)}
aa, _ := core.ComplexFromArray(az, 2, 2)
ba, _ := core.ComplexFromArray(bz, 2, 2)
a := FromArray(aa, true)
b := FromArray(ba, true)
build := func() (*Tensor, error) {
m, err := a.MatMul(b)
if err != nil {
return nil, err
}
r, err := m.Real()
if err != nil {
return nil, err
}
sq, err := r.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
a.ZeroGrad()
b.ZeroGrad()
loss, err := build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
const h = 1e-6
check := func(base *Tensor, vals []complex128, name string) {
t.Helper()
for j := range len(vals) {
perturb := func(dr, di float64) float64 {
ws := append([]complex128(nil), vals...)
ws[j] = complex(real(vals[j])+dr, imag(vals[j])+di)
base.ReplaceWith(mustRecoverComplex(ws, 2, 2))
l, err := build()
if err != nil {
t.Fatal(err)
}
return totalLoss(l)
}
dr := (perturb(h, 0) - perturb(-h, 0)) / (2 * h)
di := (perturb(0, h) - perturb(0, -h)) / (2 * h)
want := complex(0.5*dr, 0.5*di)
if got := base.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
t.Errorf("complex matmul %s grad[%d] = %g, want %g", name, j, got, want)
}
base.ReplaceWith(mustRecoverComplex(vals, 2, 2))
}
}
check(a, az, "a")
check(b, bz, "b")
}
// TestConcatMixedDtypeBackward checks the join's backward on a real and
// a complex side at once: the real side receives 2 Re g through the
// narrowing, the complex side its Wirtinger gradient.
func TestConcatMixedDtypeBackward(t *testing.T) {
r, _ := FromFloat64s([]float64{1, -2, 3}, true, 3)
zr := []complex128{complex(0.5, 1), complex(-1.5, 0.25)}
za, _ := core.ComplexFromArray(zr, 2)
z := FromArray(za, true)
build := func() (*Tensor, error) {
c, err := r.Concat(z, 0)
if err != nil {
return nil, err
}
re, err := c.Real()
if err != nil {
return nil, err
}
sq, err := re.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
r.ZeroGrad()
z.ZeroGrad()
loss, err := build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
grads := flatFloats(r.Grad())
orig := flatFloats(r.Data())
for i := range orig {
const h = 1e-6
plus := append([]float64(nil), orig...)
plus[i] += h
r.ReplaceWith(mustRecover(plus, 3))
lp, _ := build()
minus := append([]float64(nil), orig...)
minus[i] -= h
r.ReplaceWith(mustRecover(minus, 3))
lm, _ := build()
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
t.Errorf("concat real side grad[%d] = %g, central %g", i, grads[i], num)
}
r.ReplaceWith(mustRecover(orig, 3))
}
for j := range 2 {
perturb := func(dr, di float64) float64 {
ws := append([]complex128(nil), zr...)
ws[j] = complex(real(zr[j])+dr, imag(zr[j])+di)
z.ReplaceWith(mustRecoverComplex(ws, 2))
l, _ := build()
return totalLoss(l)
}
dr := (perturb(1e-6, 0) - perturb(-1e-6, 0)) / (2e-6)
di := (perturb(0, 1e-6) - perturb(0, -1e-6)) / (2e-6)
want := complex(0.5*dr, 0.5*di)
if got := z.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
t.Errorf("concat complex side grad[%d] = %g, want %g", j, got, want)
}
z.ReplaceWith(mustRecoverComplex(zr, 2))
}
}
// TestTransposeAxesBackward reverses an axis permutation under a
// non-linear loss and checks the gradient against central differences.
func TestTransposeAxesBackward(t *testing.T) {
x, _ := FromFloat64s([]float64{1, 2, 3, 4, 5, 6, 7, 8}, true, 2, 2, 2)
build := func() (*Tensor, error) {
p, err := x.TransposeAxes(2, 0, 1)
if err != nil {
return nil, err
}
sq, err := p.Abs2()
if err != nil {
return nil, err
}
w, err := sq.Mul(p)
if err != nil {
return nil, err
}
return w.Sum()
}
probeCheckCentral(t, x, "transposeaxes", build, 1e-5)
}
// TestUnaryKernelChainBackward pushes Sqrt, Log, Sigmoid, Scale, Neg,
// Mean and Pow through one graph and checks the committed gradient.
func TestUnaryKernelChainBackward(t *testing.T) {
x, _ := FromFloat64s([]float64{0.5, 1.5, 2.5, 3.5}, true, 4)
build := func() (*Tensor, error) {
s, err := x.Sqrt()
if err != nil {
return nil, err
}
l, err := s.Log()
if err != nil {
return nil, err
}
sg, err := x.Sigmoid()
if err != nil {
return nil, err
}
m, err := l.Mul(sg)
if err != nil {
return nil, err
}
sc, err := m.Scale(2.5)
if err != nil {
return nil, err
}
ng, err := sc.Neg()
if err != nil {
return nil, err
}
mn, err := ng.Mean()
if err != nil {
return nil, err
}
return mn.Pow(2)
}
probeCheckCentral(t, x, "unary-chain", build, 1e-4)
}
// TestBackwardAccumulatesUntilZeroGrad pins the accumulation contract
// end to end: two Backward calls double the gradient, ZeroGrad resets.
func TestBackwardAccumulatesUntilZeroGrad(t *testing.T) {
x, _ := FromFloat64s([]float64{2}, true, 1)
build := func() (*Tensor, error) {
s, err := x.Pow(2)
if err != nil {
return nil, err
}
return s.Sum()
}
l1, err := build()
if err != nil {
t.Fatal(err)
}
if err := l1.Backward(); err != nil {
t.Fatal(err)
}
l2, err := build()
if err != nil {
t.Fatal(err)
}
if err := l2.Backward(); err != nil {
t.Fatal(err)
}
if got := x.Grad().FloatAt(0); got != 8 {
t.Fatalf("accumulated g = %g, want 8", got)
}
x.ZeroGrad()
l3, err := build()
if err != nil {
t.Fatal(err)
}
if err := l3.Backward(); err != nil {
t.Fatal(err)
}
if got := x.Grad().FloatAt(0); got != 4 {
t.Fatalf("post-zero g = %g, want 4", got)
}
}
// TestFloat32LeafKeepsGradientDtype pins the dtype contract on a
// float32 leaf: the gradient narrows to the leaf's width.
func TestFloat32LeafKeepsGradientDtype(t *testing.T) {
v := []float32{1.5, -2.5}
a, _ := core.FromFloat32Slice(v, 2)
x := FromArray(a, true)
y, err := x.Pow(2)
if err != nil {
t.Fatal(err)
}
l, err := y.Sum()
if err != nil {
t.Fatal(err)
}
if err := l.Backward(); err != nil {
t.Fatal(err)
}
g := x.Grad()
if g.Dtype() != core.Float32 {
t.Fatalf("grad dtype = %s, want float32", g.Dtype())
}
for i := range 2 {
want := 2 * float64(v[i])
if got := g.FloatAt(i); math.Abs(got-want) > 1e-6 {
t.Fatalf("g[%d] = %g, want %g", i, got, want)
}
}
}
// TestBroadcastToBackwardMatchesCentralDifferences and its SumAxis
// neighbour cover the reduction/expansion pair in isolation.
func TestBroadcastToBackwardMatchesCentralDifferences(t *testing.T) {
x, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
build := func() (*Tensor, error) {
bc, err := x.BroadcastTo(2, 2, 3)
if err != nil {
return nil, err
}
sq, err := bc.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
probeCheckCentral(t, x, "broadcast", build, 1e-5)
}
// TestSumAxisBackwardMatchesCentralDifferences reduces the middle axis
// with a non-scalar loss, so the backward must scatter into both rows.
func TestSumAxisBackwardMatchesCentralDifferences(t *testing.T) {
x, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
build := func() (*Tensor, error) {
sq, err := x.Abs2()
if err != nil {
return nil, err
}
return sq.SumAxis(1)
}
probeCheckCentral(t, x, "sumaxis", build, 1e-5)
}
// TestMatMulTanhBackwardMatchesCentralDifferences checks the matmul
// product rule under a tanh on top, non-scalar loss included.
func TestMatMulTanhBackwardMatchesCentralDifferences(t *testing.T) {
a, _ := FromFloat64s([]float64{1, -2, 3, -4}, true, 2, 2)
w, _ := FromFloat64s([]float64{0.5, -0.25, 1, -2}, true, 2, 2)
build := func() (*Tensor, error) {
mm, err := a.MatMul(w)
if err != nil {
return nil, err
}
return mm.Tanh()
}
a.ZeroGrad()
w.ZeroGrad()
loss, err := build()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
check := func(base *Tensor, name string) {
t.Helper()
grads := flatFloats(base.Grad())
orig := flatFloats(base.Data())
sh := base.Data().Shape()
for i := range orig {
const h = 1e-6
plus := append([]float64(nil), orig...)
plus[i] += h
base.ReplaceWith(mustRecover(plus, sh...))
lp, _ := build()
minus := append([]float64(nil), orig...)
minus[i] -= h
base.ReplaceWith(mustRecover(minus, sh...))
lm, _ := build()
var num float64
for j := range lp.Data().Len() {
num += lp.Data().FloatAt(j) - lm.Data().FloatAt(j)
}
num /= 2 * h
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
t.Errorf("%s[%d] backward %g, central %g", name, i, grads[i], num)
}
base.ReplaceWith(mustRecover(orig, sh...))
}
}
check(a, "a")
check(w, "w")
}
// TestNewtonCGSolvesLeastSquares pins the truncated Newton method on a
// small full-rank least-squares problem whose solution is A^-1 b.
func TestNewtonCGSolvesLeastSquares(t *testing.T) {
ar := []float64{2, 0.5, 1, 3}
br := []float64{1, -1}
A, _ := FromFloat64s(ar, false, 2, 2)
B, _ := FromFloat64s(br, false, 2)
f := func(x *Tensor) (*Tensor, error) {
ax, err := A.MatMul(x)
if err != nil {
return nil, err
}
d, err := ax.Sub(B)
if err != nil {
return nil, err
}
s, err := d.Abs2()
if err != nil {
return nil, err
}
return s.Sum()
}
x0 := mustRecover([]float64{0, 0}, 2)
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{})
if err != nil {
t.Fatal(err)
}
det := 2*3 - 0.5*1
wantX := (3*1 - 0.5*(-1)) / det
wantY := (2*(-1) - 1*1) / det
if math.Abs(x.FloatAt(0)-wantX) > 1e-5 || math.Abs(x.FloatAt(1)-wantY) > 1e-5 {
t.Fatalf("NewtonCG point (%g, %g), want (%g, %g)", x.FloatAt(0), x.FloatAt(1), wantX, wantY)
}
if fv > 1e-10 {
t.Fatalf("NewtonCG f = %g, want ~0", fv)
}
}
+165
View File
@@ -0,0 +1,165 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// MatMulBatched multiplies stacked 3-D matrices batch-wise:
// (N, M, K) · (N, K, P) gives (N, M, P). The backward runs the classic
// product rule inside every batch slot, dA·Bᵀ and Aᵀ·dB, so batched
// sequence models can one day drop their per-slice fan-out without
// leaving the graph.
func (t *Tensor) MatMulBatched(u *Tensor) (*Tensor, error) {
if err := t.checkFloat("MatMulBatched"); err != nil {
return nil, err
}
if err := u.checkFloat("MatMulBatched"); err != nil {
return nil, err
}
ta, tu := t.data.Shape(), u.data.Shape()
if len(ta) != 3 || len(tu) != 3 {
return nil, errf("MatMulBatched: needs rank-3 operands, got %s and %s",
prettyShape(ta), prettyShape(tu))
}
if ta[0] != tu[0] || ta[2] != tu[1] {
return nil, errf("MatMulBatched: batch or inner dimension mismatch for %s · %s",
prettyShape(ta), prettyShape(tu))
}
n := ta[0]
slices := make([]*core.Array, n)
// The dtype follows the promotion ladder even for an empty batch,
// where no product runs to derive it: an Int output for float
// inputs would leak the zero value.
dt := t.data.Dtype()
if u.data.Dtype() == core.Complex || dt == core.Complex {
dt = core.Complex
} else if u.data.Dtype() == core.Float && dt != core.Float {
dt = core.Float
}
for i := range n {
aSlot, _ := core.Slice(t.data, 0, i, i+1)
bSlot, _ := core.Slice(u.data, 0, i, i+1)
aMat, _ := core.Reshape(aSlot, ta[1], ta[2])
bMat, _ := core.Reshape(bSlot, tu[1], tu[2])
prod, err := core.MatMul2D(aMat, bMat)
if err != nil {
return nil, err
}
slices[i] = prod
dt = prod.Dtype()
}
out := zeros(dt, []int{n, ta[1], tu[2]})
slot := ta[1] * tu[2]
// Each slot lands in a fresh contiguous product, so the scatter is
// a raw slice move per batch row, widened nowhere: out carries the
// products' own dtype.
for i := range n {
switch dt {
case core.Float32:
copy(out.RawFloat32s()[i*slot:(i+1)*slot], slices[i].RawFloat32s())
case core.Float:
copy(out.RawFloats()[i*slot:(i+1)*slot], slices[i].RawFloats())
default:
for j := range slot {
out.SetFloatAt(i*slot+j, slices[i].FloatAt(j))
}
}
}
at, au := t.data, u.data
return binaryResult("MatMulBatched", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
da := gradSlot{arr: ar.borrowGrad(core.Float, ta), sh: ta}
db := gradSlot{arr: ar.borrowGrad(core.Float, tu), sh: tu}
m, k, p := ta[1], ta[2], tu[2]
slotLen := m * p
for i := range n {
gMat, err := core.Reshape(
mustWindow(g.arr, i*slotLen, slotLen), m, p)
if err != nil {
return err
}
bMat, err := windowMatrix(au, i, k, p)
if err != nil {
return err
}
bT := core.Transpose(bMat)
daPart, err := core.MatMul2D(gMat, bT)
if err != nil {
return err
}
copyInto(da.arr, daPart, i*m*k)
aMat, err := windowMatrix(at, i, m, k)
if err != nil {
return err
}
aT := core.Transpose(aMat)
dbPart, err := core.MatMul2D(aT, gMat)
if err != nil {
return err
}
copyInto(db.arr, dbPart, i*k*p)
}
dst[0], dst[1] = da, db
return nil
}), nil
}
// mustWindow flattens row-major window [from, from+len) into an
// (rows, cols) matrix view materialisation. The window hands its slice
// to FloatsFromArray, which takes ownership, so contiguous gradients
// copy nothing at all.
func mustWindow(g *core.Array, from, length int) *core.Array {
if !g.Strided() && g.Dtype() == core.Float {
arr, _ := core.FloatsFromArray(g.RawFloats()[from:from+length], length)
return arr
}
vals := make([]float64, length)
for i := range vals {
vals[i] = g.FloatAt(from + i)
}
arr, _ := core.FromFloats(vals, length)
return arr
}
// windowMatrix reads batch slot n as an (r, c) float64 matrix, the
// backward's native arithmetic domain. A contiguous float64 operand is
// aliased rather than copied; anything else is read through the
// widening accessor.
func windowMatrix(a *core.Array, n, r, c int) (*core.Array, error) {
base := n * r * c
if !a.Strided() && a.Dtype() == core.Float {
arr, err := core.FloatsFromArray(a.RawFloats()[base:base+r*c], r, c)
return arr, err
}
vals := make([]float64, r*c)
for i := range vals {
vals[i] = a.FloatAt(base + i)
}
return core.FromFloats(vals, r, c)
}
// copyInto writes src's elements at dst's flat offset. The
// destination is a fresh contiguous float64 accumulator; a contiguous
// float64 source moves with one copy, a float32 one widens in place.
func copyInto(dst, src *core.Array, offset int) {
switch {
case !src.Strided() && src.Dtype() == core.Float:
copy(dst.RawFloats()[offset:], src.RawFloats())
case !src.Strided() && src.Dtype() == core.Float32:
// Bounded by the source: the destination tail runs on to the
// end of the accumulator, which is longer for every batch but
// the last.
ss, ds := src.RawFloat32s(), dst.RawFloats()[offset:]
for i := range src.Len() {
ds[i] = float64(ss[i])
}
default:
for i := range src.Len() {
dst.SetFloatAt(offset+i, src.FloatAt(i))
}
}
}
+188
View File
@@ -0,0 +1,188 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
func TestMatMulBatchedForward(t *testing.T) {
a, _ := core.FromFloats([]float64{
1, 0,
0, 1,
3, 4,
5, 6,
}, 2, 2, 2)
b, _ := core.FromFloats([]float64{
1, 1,
1, 0,
2, 0,
0, 2,
}, 2, 2, 2)
out, err := FromArray(a, false).MatMulBatched(FromArray(b, false))
if err != nil {
t.Fatal(err)
}
want := []float64{1, 1, 1, 0, 6, 8, 10, 12}
for i := range want {
if g := out.Data().FloatAt(i); g != want[i] {
t.Fatalf("slot %d = %v, want %v", i, g, want[i])
}
}
// Rank and batch mismatches error loudly.
flat, _ := core.Reshape(a, 8)
if _, err := FromArray(flat, false).MatMulBatched(FromArray(b, false)); err == nil {
t.Fatal("rank-2 operand accepted")
}
c, _ := core.FromFloats(make([]float64, 4), 1, 2, 2)
if _, err := FromArray(a, false).MatMulBatched(FromArray(c, false)); err == nil {
t.Fatal("batch-size mismatch accepted")
}
d, _ := core.FromFloats(make([]float64, 12), 2, 3, 2)
if _, err := FromArray(a, false).MatMulBatched(FromArray(d, false)); err == nil {
t.Fatal("inner-dimension mismatch accepted")
}
}
// TestMatMulBatchedGradients finite-difference checks both operands on
// a weighted sum objective so every batch slot earns its own weight.
func TestMatMulBatchedGradients(t *testing.T) {
aVal := []float64{0.5, -1, 2, 0.25, 1.5, -0.5, 0.75, 1.25, -0.25, 0.5, -1.5, 2}
bVal := []float64{1, -0.25, 0.75, 2, -1, 0.125, 0.5, -2}
weight := sweepPattern(12) // covers (2, 3, 2) outputs
aArr, _ := core.FromFloats(aVal, 2, 3, 2)
bArr, _ := core.FromFloats(bVal, 2, 2, 2)
mArr, _ := core.FromFloats(weight, 2, 3, 2)
at := FromArray(aArr, true)
bt := FromArray(bArr, true)
out, err := at.MatMulBatched(bt)
if err != nil {
t.Fatal(err)
}
scaled, err := out.Mul(FromArray(mArr, false))
if err != nil {
t.Fatal(err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
objective := func(av, bv []float64) float64 {
x, _ := core.FromFloats(av, 2, 3, 2)
y, _ := core.FromFloats(bv, 2, 2, 2)
o, rerr := FromArray(x, false).MatMulBatched(FromArray(y, false))
if rerr != nil {
return math.NaN()
}
total := 0.0
for i := range weight {
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
}
return total
}
checkSpan(t, at.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(flatten(v), bVal)
}, aArr))
checkSpan(t, bt.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(aVal, flatten(v))
}, bArr))
}
func flatten(v *core.Array) []float64 {
out := make([]float64, v.Len())
for i := range v.Len() {
out[i] = v.FloatAt(i)
}
return out
}
// TestMatMulBatchedGradientsFloat32 runs the weighted-sum
// finite-difference check on float32 operands: the forward rounds to
// float32 while the backward widens the accessors and answers float64
// gradients, and both agree with the float64 reference.
func TestMatMulBatchedGradientsFloat32(t *testing.T) {
aVal := []float32{0.5, -1, 2, 0.25, 1.5, -0.5, 0.75, 1.25, -0.25, 0.5, -1.5, 2}
bVal := []float32{1, -0.25, 0.75, 2, -1, 0.125, 0.5, -2}
weight := sweepPattern(12)
aArr, err := core.FromFloat32s(aVal, 2, 3, 2)
if err != nil {
t.Fatal(err)
}
bArr, err := core.FromFloat32s(bVal, 2, 2, 2)
if err != nil {
t.Fatal(err)
}
mArr, err := core.FromFloats(weight, 2, 3, 2)
if err != nil {
t.Fatal(err)
}
at := FromArray(aArr, true)
bt := FromArray(bArr, true)
out, err := at.MatMulBatched(bt)
if err != nil {
t.Fatal(err)
}
scaled, err := out.Mul(FromArray(mArr, false))
if err != nil {
t.Fatal(err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
if at.Grad().Dtype() != core.Float32 || bt.Grad().Dtype() != core.Float32 {
t.Fatalf("float32 leaves carry %s and %s gradients, want float32",
at.Grad().Dtype(), bt.Grad().Dtype())
}
// The reference differentiates the same batched product evaluated
// in float64 over the identical operand values.
objective := func(av, bv []float64) float64 {
x, _ := core.FromFloats(av, 2, 3, 2)
y, _ := core.FromFloats(bv, 2, 2, 2)
o, rerr := FromArray(x, false).MatMulBatched(FromArray(y, false))
if rerr != nil {
return math.NaN()
}
total := 0.0
for i := range weight {
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
}
return total
}
aRef, _ := core.FromFloats(widen32(aVal), 2, 3, 2)
bRef, _ := core.FromFloats(widen32(bVal), 2, 2, 2)
checkSpan(t, at.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(flatten(v), widen32(bVal))
}, aRef))
checkSpan(t, bt.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(widen32(aVal), flatten(v))
}, bRef))
}
// widen32 widens a float32 slice exactly, the view the backward's own
// accessors read.
func widen32(v []float32) []float64 {
out := make([]float64, len(v))
for i, x := range v {
out[i] = float64(x)
}
return out
}
+292
View File
@@ -0,0 +1,292 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The benchmarks below pin the costs the graph machinery adds around
// the kernels: one deep chain of small tensors (per-node tape cost), one
// wide element-wise graph (per-edge accumulation cost), the mid-size
// sweeps whose per-element work decides their parallel policy, and the
// L2-norm axis backward. Inputs are fixed literals, so a run is
// deterministic.
// benchLit builds a leaf of n elements from fixed literals, keeping the
// arithmetic well inside the domain of every op used here.
func benchLit(b *testing.B, seed, n int) *Tensor {
b.Helper()
v := make([]float64, n)
for i := range v {
v[i] = 0.25 + float64((i*7+seed)%13)*0.125
}
a, err := core.FromFloats(v, n)
if err != nil {
b.Fatal(err)
}
return FromArray(a, true)
}
// mustReshape reshapes an array for a benchmark fixture.
func mustReshape(b *testing.B, a *core.Array, shape ...int) *core.Array {
b.Helper()
out, err := core.Reshape(a, shape...)
if err != nil {
b.Fatal(err)
}
return out
}
// deepChain builds a chain of n element-wise nodes over x and w and
// reduces it to a scalar, the shape a training loop's tape has.
func deepChain(x, w *Tensor, n int) (*Tensor, error) {
h := x
for i := range n {
var err error
switch i % 4 {
case 0:
h, err = h.Add(w)
case 1:
h, err = h.Mul(w)
case 2:
h, err = h.Tanh()
default:
h, err = h.Scale(0.25)
}
if err != nil {
return nil, err
}
}
return h.Sum()
}
// wideFan multiplies x by n independent leaves and sums the products,
// so the backward folds n contributions into x's gradient.
func wideFan(x *Tensor, leaves []*Tensor) (*Tensor, error) {
acc, err := x.Mul(leaves[0])
if err != nil {
return nil, err
}
for _, l := range leaves[1:] {
p, err := x.Mul(l)
if err != nil {
return nil, err
}
if acc, err = acc.Add(p); err != nil {
return nil, err
}
}
return acc.Sum()
}
// BenchmarkTapeDeepForward measures the forward pass alone: one node
// per element-wise op over an 8-element tensor.
func BenchmarkTapeDeepForward(b *testing.B) {
x, w := benchLit(b, 1, 8), benchLit(b, 2, 8)
b.ReportAllocs()
for b.Loop() {
s, err := deepChain(x, w, 128)
if err != nil {
b.Fatal(err)
}
if s.Data().Len() != 1 {
b.Fatal("unexpected shape")
}
}
}
// BenchmarkTapeDeepBackward measures the same chain with the reverse
// sweep, where every node reads its gradient and folds into the two
// shared leaves.
func BenchmarkTapeDeepBackward(b *testing.B) {
x, w := benchLit(b, 1, 8), benchLit(b, 2, 8)
b.ReportAllocs()
for b.Loop() {
s, err := deepChain(x, w, 128)
if err != nil {
b.Fatal(err)
}
x.ZeroGrad()
w.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkTapeWideForward builds a 64-way fan over x, one node per
// leaf, and reduces it.
func BenchmarkTapeWideForward(b *testing.B) {
x := benchLit(b, 3, 16)
leaves := make([]*Tensor, 64)
for i := range leaves {
leaves[i] = benchLit(b, 10+i, 16)
}
b.ReportAllocs()
for b.Loop() {
s, err := wideFan(x, leaves)
if err != nil {
b.Fatal(err)
}
if s.Data().Len() != 1 {
b.Fatal("unexpected shape")
}
}
}
// BenchmarkTapeWideBackward runs the same 64-way fan with the reverse
// sweep: 64 edges fold into x's gradient through the reduction tree.
func BenchmarkTapeWideBackward(b *testing.B) {
x := benchLit(b, 3, 16)
leaves := make([]*Tensor, 64)
for i := range leaves {
leaves[i] = benchLit(b, 10+i, 16)
}
b.ReportAllocs()
for b.Loop() {
s, err := wideFan(x, leaves)
if err != nil {
b.Fatal(err)
}
x.ZeroGrad()
for _, l := range leaves {
l.ZeroGrad()
}
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// midElems is the sweep size the transcendental benchmarks use: below
// the element-wise floor of 1024 per worker on a 32-worker machine, so
// a sweep of this size runs on the calling goroutine under that policy
// and splits under a floor scaled to its per-element cost.
const midElems = 20000
// BenchmarkPowForwardMid measures the integer-exponent power over a
// mid-size sweep, one math.Pow per element.
func BenchmarkPowForwardMid(b *testing.B) {
x := benchLit(b, 5, midElems)
b.ReportAllocs()
for b.Loop() {
y, err := x.Pow(3)
if err != nil {
b.Fatal(err)
}
if y.Data().Len() != midElems {
b.Fatal("unexpected shape")
}
}
}
// BenchmarkPowBackwardMid measures the power backward over the same
// size: one Pow and one multiply per element, plus the reduction's
// fill.
func BenchmarkPowBackwardMid(b *testing.B) {
x := benchLit(b, 5, midElems)
y, err := x.Pow(3)
if err != nil {
b.Fatal(err)
}
s, err := y.Sum()
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
x.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkSqrtBackwardMid measures the square-root backward over a
// mid-size sweep, one divide per element.
func BenchmarkSqrtBackwardMid(b *testing.B) {
x := benchLit(b, 6, midElems)
y, err := x.Sqrt()
if err != nil {
b.Fatal(err)
}
s, err := y.Sum()
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
x.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkMeanAxisBackwardMid measures the axis-mean backward, whose
// first stage divides every element of the incoming gradient.
func BenchmarkMeanAxisBackwardMid(b *testing.B) {
x := benchLit(b, 7, midElems)
xt := FromArray(mustReshape(b, x.Data(), 200, 100), true)
y, err := xt.MeanAxis(0)
if err != nil {
b.Fatal(err)
}
s, err := y.Sum()
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
xt.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkSumBackwardWide measures a reduction over a large operand:
// the backward fills the operand's shape with one value, and the pass
// commits a hundred-thousand-element gradient into the leaf.
func BenchmarkSumBackwardWide(b *testing.B) {
x := benchLit(b, 8, 100000)
s, err := x.Sum()
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
x.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkL2NormAxisBackward measures the norm backward over a
// (2000×100) tensor reduced along the leading axis: 100 lines of 2000
// elements each, long enough for the line sweep to dominate the
// allocation of the output.
func BenchmarkL2NormAxisBackward(b *testing.B) {
x := benchLit(b, 9, 200000)
xt := FromArray(mustReshape(b, x.Data(), 2000, 100), true)
y, err := xt.L2NormAxis(0)
if err != nil {
b.Fatal(err)
}
s, err := y.Sum()
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
xt.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
+112
View File
@@ -0,0 +1,112 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"testing"
core "sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Backward benchmarks guard the tape overhead around the kernels: the
// two-layer graph is the smallest shape where per-node costs and the
// matmul backward both show.
func benchVals(b *testing.B, seed, n int) []float64 {
b.Helper()
v := make([]float64, n)
for i := range v {
v[i] = float64(i%13)*float64(seed%3)*0.25 + float64(i%5) - 2
}
return v
}
func benchTensor(b *testing.B, seed int, shape ...int) *Tensor {
b.Helper()
n := 1
for _, d := range shape {
n *= d
}
a, err := core.FromFloats(benchVals(b, seed, n), shape...)
if err != nil {
b.Fatal(err)
}
return FromArray(a, true)
}
// BenchmarkBackwardTwoLayer runs forward and backward over
// (32×64)·(64×32), then tanh, then ·(32×10), then sum.
func BenchmarkBackwardTwoLayer(b *testing.B) {
x := benchTensor(b, 1, 32, 64)
w1 := benchTensor(b, 2, 64, 32)
w2 := benchTensor(b, 3, 32, 10)
b.ReportAllocs()
for b.Loop() {
h, err := x.MatMul(w1)
if err != nil {
b.Fatal(err)
}
t, err := h.Tanh()
if err != nil {
b.Fatal(err)
}
y, err := t.MatMul(w2)
if err != nil {
b.Fatal(err)
}
s, err := y.Sum()
if err != nil {
b.Fatal(err)
}
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkForwardOnly isolates the graph construction from the
// backward sweep.
func BenchmarkForwardOnly(b *testing.B) {
x := benchTensor(b, 1, 32, 64)
w1 := benchTensor(b, 2, 64, 32)
w2 := benchTensor(b, 3, 32, 10)
b.ReportAllocs()
for b.Loop() {
h, err := x.MatMul(w1)
if err != nil {
b.Fatal(err)
}
t, err := h.Tanh()
if err != nil {
b.Fatal(err)
}
if _, err := t.MatMul(w2); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkGradMatMulBackward isolates one MatMul node's backward
// sweep on (128×128) operands.
func BenchmarkGradMatMulBackward(b *testing.B) {
a := benchTensor(b, 4, 128, 128)
c := benchTensor(b, 5, 128, 128)
y, err := a.MatMul(c)
if err != nil {
b.Fatal(err)
}
s, err := y.Sum()
if err != nil {
b.Fatal(err)
}
b.ResetTimer()
b.ReportAllocs()
for b.Loop() {
a.ZeroGrad()
c.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
+145
View File
@@ -0,0 +1,145 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// Regression pins for the broadcast and shape backwards: a broadcast
// gradient must collapse to its source's shape, and every widened op
// must hand each operand its own correctly shaped gradient buffer.
// TestBroadcastRank1Gradient pins the rank-1 broadcast backward: the
// gradient of a size-1 source broadcast to length m must collapse back
// to a single sum, not arrive with the broadcast shape.
func TestBroadcastRank1Gradient(t *testing.T) {
x, err := core.FromFloats([]float64{0}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
y, err := xt.BroadcastTo(3)
if err != nil {
t.Fatalf("BroadcastTo: %v", err)
}
w, err := core.FromFloats([]float64{1, 2, 3}, 3)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
wt := FromArray(w, false)
prod, err := y.Mul(wt)
if err != nil {
t.Fatalf("Mul: %v", err)
}
loss, err := prod.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if g.NDim() != 1 || g.Shape()[0] != 1 {
t.Fatalf("gradient shape %v, want [1]", g.Shape())
}
if math.Abs(g.FloatAt(0)-6) > 1e-12 {
t.Fatalf("gradient = %g, want 6", g.FloatAt(0))
}
}
// TestBroadcastRank1ToMatrix pins the (1,) to (m, n) broadcast backward
// against central differences.
func TestBroadcastRank1ToMatrix(t *testing.T) {
x, err := core.FromFloats([]float64{2}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
y, err := xt.BroadcastTo(2, 3)
if err != nil {
t.Fatalf("BroadcastTo: %v", err)
}
w, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
wt := FromArray(w, false)
prod, err := y.Mul(wt)
if err != nil {
t.Fatalf("Mul: %v", err)
}
loss, err := prod.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if g.NDim() != 1 || g.Shape()[0] != 1 {
t.Fatalf("gradient shape %v, want [1]", g.Shape())
}
if math.Abs(g.FloatAt(0)-21) > 1e-12 {
t.Fatalf("gradient = %g, want 21", g.FloatAt(0))
}
}
// TestL2NormAxisEmptyDim pins the backward against the integer division
// by zero an empty reduced dimension used to hit.
func TestL2NormAxisEmptyDim(t *testing.T) {
x, err := core.FromFloats([]float64{}, 3, 0)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
n, err := xt.L2NormAxis(1)
if err != nil {
t.Fatalf("L2NormAxis: %v", err)
}
loss, err := n.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if g.Len() != 0 {
t.Fatalf("gradient length %d, want 0", g.Len())
}
}
// TestShapeOpsRejectNonFloat pins the dtype contract on the shape ops
// that used to record graph nodes without validating the dtype.
func TestShapeOpsRejectNonFloat(t *testing.T) {
i, err := core.FromInts([]int64{1, 2, 3, 4}, 2, 2)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
it := FromArray(i, true)
if _, err := it.Transpose(); err == nil {
t.Error("Transpose accepted an int tensor")
}
if _, err := it.Squeeze(0); err == nil {
t.Error("Squeeze accepted an int tensor")
}
if _, err := it.Unsqueeze(0); err == nil {
t.Error("Unsqueeze accepted an int tensor")
}
if _, err := it.Reshape(4); err == nil {
t.Error("Reshape accepted an int tensor")
}
if _, err := it.TransposeAxes(1, 0); err == nil {
t.Error("TransposeAxes accepted an int tensor")
}
if _, err := it.Floor(); err == nil {
t.Error("Floor accepted an int tensor")
}
if _, err := it.BroadcastTo(2, 3); err == nil {
t.Error("BroadcastTo accepted an int tensor")
}
}
+503
View File
@@ -0,0 +1,503 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Complex differentiation. The graph accepts complex128
// tensors alongside float64/float32, with the Wirtinger convention the
// optimiser ecosystem settled on: Backward seeds a REAL scalar loss
// (a complex output is rejected with an error telling the caller to
// reduce first), and the gradient a complex leaf accumulates is
// ∂L/∂z̄, the direction gradient descent steps along. Under that
// convention the adjoint of a holomorphic op y = f(z) is
// dz += g·conj(f′(z)), so every conjugation below sits exactly where
// the calculus puts it.
//
// A real tensor inside a complex graph narrows the incoming complex
// gradient by 2·Re: for a real variable x, dL/dx = 2·Re(∂L/∂x̄), and
// the factor also cancels the ½ the Real backward contributes, so
// mixed graphs compose exactly.
// checkDiff is checkFloat plus complex: the ops that can differentiate
// complex inputs validate with it.
func (t *Tensor) checkDiff(name string) error {
switch t.data.Dtype() {
case core.Float, core.Float32, core.Complex:
return nil
}
return errf("autograd: %s needs a float, float32 or complex tensor, got %s", name, t.data.Dtype())
}
// isComplexArr reports whether a holds complex128 data.
func isComplexArr(a *core.Array) bool { return a.Dtype() == core.Complex }
// eitherComplex reports whether either operand is complex.
func eitherComplex(a, b *core.Array) bool { return isComplexArr(a) || isComplexArr(b) }
// conjArray returns the element-wise conjugate. Real arrays come back
// unchanged (their conjugate is themselves), so mixed-dtype adjoints
// can call it unconditionally.
func conjArray(a *core.Array) *core.Array {
if !isComplexArr(a) {
return a
}
out := zeros(core.Complex, a.Shape())
cs := out.RawComplexes()
if a.Strided() {
for i := range cs {
cs[i] = conj(a.ComplexAt(i))
}
return out
}
as := a.RawComplexes()
for i := range cs {
z := as[i]
cs[i] = complex(real(z), -imag(z))
}
return out
}
func conj(z complex128) complex128 { return complex(real(z), -imag(z)) }
// copyElem copies one element between gradient arrays of the same
// dtype; the callers narrow the incoming gradient to the operand's
// dtype with narrowGradient before the copy, so a mixed real/complex
// pair never reaches here and a real destination never reads a complex
// payload. A complex destination reads through ComplexAt, which serves
// a strided source too.
func copyElem(dst *core.Array, di int, src *core.Array, si int) {
if dst.Dtype() == core.Complex {
dst.RawComplexes()[di] = src.ComplexAt(si)
return
}
dst.SetFloatAt(di, src.FloatAt(si))
}
// narrowGradient converts a gradient to the dtype of the tensor it
// accumulates into. Complex to real takes 2·Re (the real-tensor rule
// above); everything else routes through Astype.
func narrowGradient(g gradSlot, dt core.Dtype) (gradSlot, error) {
if g.arr.Dtype() == dt {
return g, nil
}
sh := g.sh
if sh == nil {
sh = g.arr.Shape()
}
if g.arr.Dtype() == core.Complex && dt != core.Complex {
out := zeros(dt, sh)
gs := g.arr.RawComplexes()
if dt == core.Float32 && !g.arr.Strided() {
os := out.RawFloat32s()
for i := range os {
os[i] = float32(2 * real(gs[i]))
}
return gradSlot{arr: out, sh: sh}, nil
}
if dt == core.Float && !g.arr.Strided() {
os := out.RawFloats()
for i := range os {
os[i] = 2 * real(gs[i])
}
return gradSlot{arr: out, sh: sh}, nil
}
// out is freshly allocated and dense, so a real destination
// takes its payload directly; the complex source keeps the
// accessor read that rebases a strided index.
switch dt {
case core.Float32:
os := out.RawFloat32s()
for i := range os {
os[i] = float32(2 * real(g.arr.ComplexAt(i)))
}
return gradSlot{arr: out, sh: sh}, nil
case core.Float:
os := out.RawFloats()
for i := range os {
os[i] = 2 * real(g.arr.ComplexAt(i))
}
return gradSlot{arr: out, sh: sh}, nil
}
for i := range g.arr.Len() {
out.SetFloatAt(i, 2*real(g.arr.ComplexAt(i)))
}
return gradSlot{arr: out, sh: sh}, nil
}
c, err := core.Astype(g.arr, dt)
if err != nil {
return gradSlot{}, err
}
return gradSlot{arr: c, sh: sh}, nil
}
// scalarComplex builds a 1-element complex array holding z.
func scalarComplex(z complex128) *core.Array {
out := zeros(core.Complex, []int{1})
out.RawComplexes()[0] = z
return out
}
// fillComplex returns a complex array shaped like a with every element
// set to z.
func fillComplex(a *core.Array, z complex128) *core.Array {
out := zeros(core.Complex, a.Shape())
cs := out.RawComplexes()
for i := range cs {
cs[i] = z
}
return out
}
// Conj returns the element-wise complex conjugate. The conjugate is
// anti-holomorphic: its ∂/∂z̄ adjoint conjugates the incoming
// gradient (dz = conj(g)), which is what makes ⟨ψ|H|ψ⟩ come out as
// Hψ rather than only its real part.
func (t *Tensor) Conj() (*Tensor, error) {
if err := t.checkDiff("Conj"); err != nil {
return nil, err
}
a := t.data
out := conjArray(t.data)
if !isComplexArr(t.data) {
// conj of a real tensor is a copy, so the graph needs its own
// node data, not the operand alias.
out = cloneReal(t.data)
}
return t.unaryResult("Conj", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
sh := gradShape(g, a)
if !isComplexArr(g.arr) {
c, err := copyGradSlot(ar, g, sh)
if err != nil {
return err
}
dst[0] = c
return nil
}
n := g.arr.Len()
da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
cs := da.arr.RawComplexes()[:n]
gs := g.arr.RawComplexes()[:n]
for i := range cs {
z := gs[i]
cs[i] = complex(real(z), -imag(z))
}
dst[0] = da
return nil
}), nil
}
// cloneReal copies a real array (the graph never aliases operands).
func cloneReal(a *core.Array) *core.Array {
out := zeros(a.Dtype(), a.Shape())
switch {
case a.Strided():
for i := range a.Len() {
out.SetFloatAt(i, a.FloatAt(i))
}
case a.Dtype() == core.Float32:
copy(out.RawFloat32s(), a.RawFloat32s())
case a.Dtype() == core.Float:
copy(out.RawFloats(), a.RawFloats())
default:
copy(out.RawInts(), a.RawInts())
}
return out
}
// Real returns the real part of each element as a float tensor. The
// complex backward halves the gradient (∂Re z/∂z̄ = ½), which the
// 2·Re narrowing at any real destination cancels exactly.
func (t *Tensor) Real() (*Tensor, error) {
if err := t.checkDiff("Real"); err != nil {
return nil, err
}
if !isComplexArr(t.data) {
// Real of a real tensor is a copy with its own storage.
out := cloneReal(t.data)
a := t.data
return t.unaryResult("Real", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
c, err := copyGradSlot(ar, g, gradShape(g, a))
if err != nil {
return err
}
dst[0] = c
return nil
}), nil
}
out := zeros(core.Float, t.data.Shape())
if t.data.Strided() {
for i := range t.data.Len() {
out.SetFloatAt(i, real(t.data.ComplexAt(i)))
}
} else {
cs := t.data.RawComplexes()
os := out.RawFloats()
for i := range os {
os[i] = real(cs[i])
}
}
// The shape is captured now: nothing may be read off the input at
// backward time, or a ReplaceWith in between would change it.
shape := t.data.Shape()
return t.unaryResult("Real", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
da := gradSlot{arr: ar.borrowGrad(core.Complex, shape), sh: shape}
cs := da.arr.RawComplexes()[:da.arr.Len()]
if g.arr.Strided() || g.arr.Dtype() != core.Float {
for i := range cs {
cs[i] = complex(g.arr.FloatAt(i)/2, 0)
}
dst[0] = da
return nil
}
gs := g.arr.RawFloats()
for i := range cs {
cs[i] = complex(gs[i]/2, 0)
}
dst[0] = da
return nil
}), nil
}
// Imag returns the imaginary part of each element as a float tensor;
// the complex backward scales by i/2 (∂Im z/∂z̄ = i/2).
func (t *Tensor) Imag() (*Tensor, error) {
if !isComplexArr(t.data) {
return nil, errf("autograd: Imag needs a complex tensor, got %s", t.data.Dtype())
}
out := zeros(core.Float, t.data.Shape())
if t.data.Strided() {
for i := range t.data.Len() {
out.SetFloatAt(i, imag(t.data.ComplexAt(i)))
}
} else {
cs := t.data.RawComplexes()
os := out.RawFloats()
for i := range os {
os[i] = imag(cs[i])
}
}
shape := t.data.Shape()
return t.unaryResult("Imag", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
da := gradSlot{arr: ar.borrowGrad(core.Complex, shape), sh: shape}
cs := da.arr.RawComplexes()[:da.arr.Len()]
if g.arr.Strided() || g.arr.Dtype() != core.Float {
for i := range cs {
cs[i] = complex(0, g.arr.FloatAt(i)/2)
}
dst[0] = da
return nil
}
gs := g.arr.RawFloats()
for i := range cs {
cs[i] = complex(0, gs[i]/2)
}
dst[0] = da
return nil
}), nil
}
// Abs2 returns |z|² of each element, a real tensor. The complex
// backward is dz = g·z (∂|z|²/∂z̄ = z); the real input path is the
// square with its 2x backward, keeping the operand's width (a float32
// input squares in float64 and stays float32, exactly as Pow
// does).
func (t *Tensor) Abs2() (*Tensor, error) {
if err := t.checkDiff("Abs2"); err != nil {
return nil, err
}
if !isComplexArr(t.data) {
return t.squareGraph()
}
a := t.data
out := zeros(core.Float, a.Shape())
if a.Strided() {
for i := range a.Len() {
z := a.ComplexAt(i)
out.SetFloatAt(i, real(z)*real(z)+imag(z)*imag(z))
}
} else {
// Bound the walk by the destination's length: a rebased view's
// payload may run longer than its element count.
as := a.RawComplexes()
os := out.RawFloats()
for i := range os {
z := as[i]
os[i] = real(z)*real(z) + imag(z)*imag(z)
}
}
return t.unaryResult("Abs2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
sh := gradShape(g, a)
da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
cs := da.arr.RawComplexes()[:da.arr.Len()]
if a.Strided() || g.arr.Strided() || g.arr.Dtype() != core.Float {
for i := range cs {
cs[i] = complex(g.arr.FloatAt(i), 0) * a.ComplexAt(i)
}
dst[0] = da
return nil
}
as, gs := a.RawComplexes(), g.arr.RawFloats()
for i := range cs {
cs[i] = complex(gs[i], 0) * as[i]
}
dst[0] = da
return nil
}), nil
}
// squareGraph is the real-input branch of Abs2: y = x², dx = 2x·g.arr.
// The output keeps the operand's width, squared in float64 and rounded
// once, exactly as Pow does, so Abs2 and Pow(2) agree on dtype
// and value for a float32 operand.
func (t *Tensor) squareGraph() (*Tensor, error) {
a := t.data
out := zeros(a.Dtype(), a.Shape())
switch {
case a.Dtype() == core.Float32 && !a.Strided():
as, os := a.RawFloat32s(), out.RawFloat32s()
for i := range os {
v := float64(as[i])
os[i] = float32(v * v)
}
case a.Strided():
for i := range out.Len() {
v := a.FloatAt(i)
out.SetFloatAt(i, v*v)
}
default:
as, os := a.RawFloats(), out.RawFloats()
for i := range os {
v := as[i]
os[i] = v * v
}
}
return t.unaryResult("Abs2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
sh := gradShape(g, a)
// dx = 2·x·g with the staged chain's rounding: the product
// forms first and the doubling multiplies it, per element.
if !a.Strided() && !g.arr.Strided() && a.Dtype() == g.arr.Dtype() && a.Len() == g.arr.Len() {
n := a.Len()
switch a.Dtype() {
case core.Float:
da := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh}
as, gs, ds := a.RawFloats()[:n], g.arr.RawFloats()[:n], da.arr.RawFloats()[:n]
for i := range ds {
ds[i] = (as[i] * gs[i]) * 2
}
dst[0] = da
return nil
case core.Float32:
da := gradSlot{arr: ar.borrowGrad(core.Float32, sh), sh: sh}
as, gs, ds := a.RawFloat32s()[:n], g.arr.RawFloat32s()[:n], da.arr.RawFloat32s()[:n]
for i := range ds {
p := float32(float64(as[i]) * float64(gs[i]))
ds[i] = float32(float64(p) * 2)
}
dst[0] = da
return nil
}
}
da, err := core.Mul(a, g.arr)
if err != nil {
return err
}
dst[0] = gradSlot{arr: core.MulI(da, 2), sh: sh}
return nil
}), nil
}
// Abs returns the absolute value of each element: complex input yields
// float magnitudes with dz = g·z/(2|z|) (zero at the origin, the
// subgradient). The real branch lives beside the other real kernels in
// tensor.go and dispatches here for complex input.
func (t *Tensor) absComplex() (*Tensor, error) {
a := t.data
out := core.Abs(a)
return t.unaryResult("Abs", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
sh := gradShape(g, a)
da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
cs := da.arr.RawComplexes()[:da.arr.Len()]
if a.Strided() || g.arr.Strided() || out.Strided() ||
g.arr.Dtype() != core.Float || out.Dtype() != core.Float {
for i := range cs {
z := a.ComplexAt(i)
m := out.FloatAt(i)
if m == 0 {
continue
}
cs[i] = complex(g.arr.FloatAt(i)/(2*m), 0) * z
}
dst[0] = da
return nil
}
as, gs, os := a.RawComplexes(), g.arr.RawFloats(), out.RawFloats()
for i := range cs {
m := os[i]
if m == 0 {
continue
}
cs[i] = complex(gs[i]/(2*m), 0) * as[i]
}
dst[0] = da
return nil
}), nil
}
// powComplexGrad builds the Wirtinger backward of y = zⁿ:
// dz = g·n·conj(z)ⁿ⁻¹, assembled by repeated conjugate multiplication
// (the exponent is a small integer; a loop beats a general power).
// sh is the shape the incoming gradient carries, or the operand's own
// on the legacy sweep path (gradShape).
func powComplexGrad(ar *gradArena, g gradSlot, a *core.Array, n int64, sh []int) *core.Array {
da := ar.borrowGrad(core.Complex, sh)
cs := da.RawComplexes()[:da.Len()]
if a.Strided() || g.arr.Strided() {
for i := range cs {
term := complex(1, 0)
for range n - 1 {
term *= conj(a.ComplexAt(i))
}
cs[i] = complex(float64(n), 0) * g.arr.ComplexAt(i) * term
}
return da
}
as, gs := a.RawComplexes(), g.arr.RawComplexes()
for i := range cs {
term := complex(1, 0)
for range n - 1 {
term *= conj(as[i])
}
cs[i] = complex(float64(n), 0) * gs[i] * term
}
return da
}
// powComplexForward raises each complex element to a non-negative
// integer power by repeated multiplication.
func powComplexForward(a *core.Array, n int64) *core.Array {
out := zeros(core.Complex, a.Shape())
cs := out.RawComplexes()
if a.Strided() {
for i := range cs {
p := complex(1, 0)
for range n {
p *= a.ComplexAt(i)
}
cs[i] = p
}
return out
}
as := a.RawComplexes()
for i := range cs {
p := complex(1, 0)
for range n {
p *= as[i]
}
cs[i] = p
}
return out
}
+539
View File
@@ -0,0 +1,539 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// buildComplex wraps a complex array as a leaf tensor; the losses
// built on it reduce to a real scalar, so Backward has a real seed.
func buildComplex(t *testing.T, vals []complex128, shape ...int) *Tensor {
t.Helper()
a, err := core.FromComplexes(vals, shape...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
return FromArray(a, true)
}
// numericComplexGrad estimates dL/dRe(z) and dL/dIm(z) by central
// differences; the Wirtinger gradient the graph reports must satisfy
// g = (dL/dRe + i·dL/dIm)/2 element-wise.
func numericComplexGrad(f func(*core.Array) float64, a *core.Array) []complex128 {
n := a.Len()
out := make([]complex128, n)
const h = 1e-6
up := make([]complex128, n)
down := make([]complex128, n)
for i := range n {
base := make([]complex128, n)
for j := range n {
base[j] = a.ComplexAt(j)
}
copy(up, base)
copy(down, base)
up[i] += complex(h, 0)
down[i] -= complex(h, 0)
au, _ := core.FromComplexes(up, a.Shape()...)
ad, _ := core.FromComplexes(down, a.Shape()...)
dRe := (f(au) - f(ad)) / (2 * h)
copy(up, base)
copy(down, base)
up[i] += complex(0, h)
down[i] -= complex(0, h)
au, _ = core.FromComplexes(up, a.Shape()...)
ad, _ = core.FromComplexes(down, a.Shape()...)
dIm := (f(au) - f(ad)) / (2 * h)
out[i] = complex(dRe/2, dIm/2)
}
return out
}
// checkAgainstNumeric compares a Wirtinger gradient with the central-
// difference reference.
func checkAgainstNumeric(t *testing.T, got *core.Array, want []complex128, tol float64) {
t.Helper()
for i, w := range want {
g := got.ComplexAt(i)
if math.Abs(real(g)-real(w)) > tol || math.Abs(imag(g)-imag(w)) > tol {
t.Fatalf("grad[%d] = %v, want %v", i, g, w)
}
}
}
// TestComplexMulGrad pins the Wirtinger adjoint of the element-wise
// product: dz = g·w̄.
func TestComplexMulGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i, -0.5 + 0.25i}, 3)
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i, 1 - 3i}, 3)
prod, err := z.Mul(w)
if err != nil {
t.Fatalf("Mul: %v", err)
}
re, err := prod.Real()
if err != nil {
t.Fatalf("Real: %v", err)
}
loss, err := re.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += real(a.ComplexAt(i) * w.Data().ComplexAt(i))
}
return s
}
want := numericComplexGrad(lossOf, z.Data())
checkAgainstNumeric(t, z.Grad(), want, 1e-8)
}
// TestComplexDivGrad pins the division adjoint da = g/b̄,
// db = −g·ā/b̄².
func TestComplexDivGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i}, 2)
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i}, 2)
q, err := z.Div(w)
if err != nil {
t.Fatalf("Div: %v", err)
}
im, err := q.Imag()
if err != nil {
t.Fatalf("Imag: %v", err)
}
loss, err := im.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += imag(a.ComplexAt(i) / w.Data().ComplexAt(i))
}
return s
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(lossOf, z.Data()), 1e-8)
checkAgainstNumeric(t, w.Grad(), numericComplexGrad(func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += imag(z.Data().ComplexAt(i) / a.ComplexAt(i))
}
return s
}, w.Data()), 1e-8)
}
// TestComplexQuantumExpectation pins the physics workhorse: the loss
// L = Re(ψ̄·(H·ψ)) with a Hermitian H, whose Wirtinger gradient is
// ψ̄-independent and equals H·ψ... evaluated against central
// differences rather than trust the algebra.
func TestComplexQuantumExpectation(t *testing.T) {
psi := buildComplex(t, []complex128{1 + 0.5i, -0.3 + 0.8i, 0.2 - 1.1i, 0.9 + 0.4i}, 4)
hDense := []complex128{
2, 0.5i, 0, -1,
-0.5i, 3, 1i, 0,
0, -1i, 1.5, 0.5,
-1, 0, 0.5, 2.5,
}
hArr, err := core.FromComplexes(hDense, 4, 4)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
h := FromArray(hArr, false)
conj, err := psi.Conj()
if err != nil {
t.Fatalf("Conj: %v", err)
}
// Row vector (1,4) times H·ψ (4,) keeps every MatMul shape legal.
bra, err := conj.Reshape(1, 4)
if err != nil {
t.Fatalf("Reshape: %v", err)
}
hpsi, err := h.MatMul(psi)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
prod, err := bra.MatMul(hpsi)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
// prod is (1,1); Real then Sum flattens to the scalar loss.
r, err := prod.Real()
if err != nil {
t.Fatalf("Real: %v", err)
}
loss, err := r.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
expectation := func(a *core.Array) float64 {
s := 0.0
for i := range 4 {
var acc complex128
for j := range 4 {
acc += hDense[i*4+j] * a.ComplexAt(j)
}
s += real(complex(real(a.ComplexAt(i)), -imag(a.ComplexAt(i))) * acc)
}
return s
}
checkAgainstNumeric(t, psi.Grad(), numericComplexGrad(expectation, psi.Data()), 1e-7)
}
// TestComplexAbs2Grad pins |z|², dz = g·z.
func TestComplexAbs2Grad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, -3 + 0.5i}, 2)
sq, err := z.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// d|z|²/dz̄ = z exactly.
for i := range 2 {
if z.Grad().ComplexAt(i) != z.Data().ComplexAt(i) {
t.Fatalf("grad[%d] = %v, want %v", i, z.Grad().ComplexAt(i), z.Data().ComplexAt(i))
}
}
}
// TestComplexAbsGrad pins the magnitude gradient dz = g·z/(2|z|).
func TestComplexAbsGrad(t *testing.T) {
z := buildComplex(t, []complex128{3 + 4i, -1 + 1i}, 2)
m, err := z.Abs()
if err != nil {
t.Fatalf("Abs: %v", err)
}
loss, err := m.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
expectation := func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += cmplxAbs(a.ComplexAt(i))
}
return s
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(expectation, z.Data()), 1e-8)
}
func cmplxAbs(z complex128) float64 { return math.Hypot(real(z), imag(z)) }
// TestComplexBackwardRejectsComplexLoss pins the real-seed contract.
func TestComplexBackwardRejectsComplexLoss(t *testing.T) {
z := buildComplex(t, []complex128{1 + 1i}, 1)
if err := z.Backward(); err == nil {
t.Fatal("Backward accepted a complex output")
}
}
// TestComplexMixedRealLeaf pins the 2·Re narrowing: a real tensor
// multiplied into a complex chain gets the true real gradient.
func TestComplexMixedRealLeaf(t *testing.T) {
xArr, err := core.FromFloats([]float64{1.5, -0.5}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
x := FromArray(xArr, true)
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i}, 2)
prod, err := x.Mul(w)
if err != nil {
t.Fatalf("Mul: %v", err)
}
sq, err := prod.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// L = Σ x²|w|², dL/dx = 2x|w|².
want := []float64{2 * 1.5 * (0.25 + 1), 2 * -0.5 * (4 + 4)}
for i, wv := range want {
if math.Abs(x.Grad().FloatAt(i)-wv) > 1e-12 {
t.Fatalf("x.grad[%d] = %g, want %g", i, x.Grad().FloatAt(i), wv)
}
}
}
// TestComplexSumAxisGrad pins the axis reduction on a complex leaf:
// the backward broadcasts the seed back over the dropped axis, and the
// Wirtinger gradient of a weighted real loss matches central
// differences.
func TestComplexSumAxisGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, -0.5 + 0.25i, 0.75 - 1.5i, 2 + 0.5i}, 2, 2)
out, err := z.SumAxis(1)
if err != nil {
t.Fatalf("SumAxis: %v", err)
}
if out.Data().NDim() != 1 || out.Data().Len() != 2 {
t.Fatalf("SumAxis shape = %s, want (2)", prettyShape(out.Data().Shape()))
}
w := buildComplex(t, []complex128{0.3 + 0.4i, -0.6 - 0.1i}, 2)
prod, err := out.Mul(w)
if err != nil {
t.Fatalf("Mul: %v", err)
}
re, err := prod.Real()
if err != nil {
t.Fatalf("Real: %v", err)
}
loss, err := re.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
s := 0.0
for j := range 2 {
var acc complex128
for k := range 2 {
acc += a.ComplexAt(j*2 + k)
}
s += real(acc * w.Data().ComplexAt(j))
}
return s
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(lossOf, z.Data()), 1e-8)
}
// TestComplexMotionOps pins gradient flow through Slice, Concat,
// Reshape and BroadcastTo on complex tensors.
func TestComplexMotionOps(t *testing.T) {
z := buildComplex(t, []complex128{1 + 1i, 2 - 1i, 3 + 2i, 4 - 3i}, 4)
sl, err := z.Slice(0, 1, 3)
if err != nil {
t.Fatalf("Slice: %v", err)
}
sq, err := sl.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// Only the sliced elements receive gradient, and for L = Σ|z|² it
// is exactly z_i.
for i := range 4 {
want := complex(0, 0)
if i == 1 || i == 2 {
want = z.Data().ComplexAt(i)
}
if z.Grad().ComplexAt(i) != want {
t.Fatalf("grad[%d] = %v, want %v", i, z.Grad().ComplexAt(i), want)
}
}
z2 := buildComplex(t, []complex128{0.5 + 0.5i}, 1)
cat, err := z.Concat(z2, 0)
if err != nil {
t.Fatalf("Concat: %v", err)
}
sq2, err := cat.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss2, err := sq2.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss2.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
if z2.Grad().ComplexAt(0) != 0.5+0.5i {
t.Fatalf("concat gradient = %v, want 0.5+0.5i", z2.Grad().ComplexAt(0))
}
}
// TestComplexSumMeanPow pins the complex reducers and integer powers.
func TestComplexSumMeanPow(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i}, 2)
m, err := z.Mean()
if err != nil {
t.Fatalf("Mean: %v", err)
}
if m.Data().ComplexAt(0) != 2+0.5i {
t.Fatalf("mean = %v, want 2+0.5i", m.Data().ComplexAt(0))
}
p, err := z.Pow(3)
if err != nil {
t.Fatalf("Pow: %v", err)
}
// (1+2i)³ = (1+2i)(1+2i)(1+2i) = -11-2i.
if p.Data().ComplexAt(0) != -11-2i {
t.Fatalf("pow = %v, want -11-2i", p.Data().ComplexAt(0))
}
sq, err := p.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
pv := complex(1, 0)
for range 3 {
pv *= a.ComplexAt(i)
}
s += real(pv)*real(pv) + imag(pv)*imag(pv)
}
return s
}, z.Data()), 1e-6)
}
// TestComplexMatMulGrad pins the 2-D complex matmul adjoint against
// central differences.
func TestComplexMatMulGrad(t *testing.T) {
aVals := []complex128{1 + 1i, 2 - 1i, 0.5 + 0i, -1 + 2i}
bVals := []complex128{0.5 - 0.5i, 1 + 1i, -0.5 + 2i, 0.25 - 0.75i}
a := buildComplex(t, aVals, 2, 2)
b := buildComplex(t, bVals, 2, 2)
y, err := a.MatMul(b)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
sq, err := y.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(av, bv []complex128) float64 {
s := 0.0
for i := range 2 {
for j := range 2 {
var acc complex128
for k := range 2 {
acc += av[i*2+k] * bv[k*2+j]
}
s += real(acc)*real(acc) + imag(acc)*imag(acc)
}
}
return s
}
checkAgainstNumeric(t, a.Grad(), numericComplexGrad(func(arr *core.Array) float64 {
return lossOf(flatComplex(arr), bVals)
}, a.Data()), 1e-7)
checkAgainstNumeric(t, b.Grad(), numericComplexGrad(func(arr *core.Array) float64 {
return lossOf(aVals, flatComplex(arr))
}, b.Data()), 1e-7)
}
func flatComplex(a *core.Array) []complex128 {
out := make([]complex128, a.Len())
for i := range out {
out[i] = a.ComplexAt(i)
}
return out
}
// TestComplexScaleGrad pins Scale on complex tensors.
func TestComplexScaleGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 1i, 2 - 1i}, 2)
s, err := z.Scale(2.5)
if err != nil {
t.Fatalf("Scale: %v", err)
}
sq, err := s.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// d(2.5²|z|²)/dz̄ = 2·2.5²·Re-parts... exact: 6.25·z.
for i := range 2 {
want := 6.25 * z.Data().ComplexAt(i)
got := z.Grad().ComplexAt(i)
if math.Abs(real(got)-real(want)) > 1e-10 || math.Abs(imag(got)-imag(want)) > 1e-10 {
t.Fatalf("grad[%d] = %v, want %v", i, got, want)
}
}
}
// TestComplexExpGrad pins the complex exponential adjoint dz = g·conj(e^z)
// against central differences, and e^z against the polar identity
// e^{x+iy} = e^x(cos y + i sin y).
func TestComplexExpGrad(t *testing.T) {
vals := []complex128{0.3 - 0.2i, -1.1 + 0.7i, 0.05 + 0i}
z := buildComplex(t, vals, 3)
e, err := z.Exp()
if err != nil {
t.Fatalf("Exp: %v", err)
}
for i, v := range vals {
want := complex(math.Exp(real(v)), 0) * complex(math.Cos(imag(v)), math.Sin(imag(v)))
got := e.Data().ComplexAt(i)
if cmplxAbs(got-want) > 1e-14 {
t.Fatalf("exp[%d] = %v, want %v", i, got, want)
}
}
sq, err := e.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
vv := a.ComplexAt(i)
ev := complex(math.Exp(real(vv)), 0) * complex(math.Cos(imag(vv)), math.Sin(imag(vv)))
s += real(ev)*real(ev) + imag(ev)*imag(ev)
}
return s
}, z.Data()), 1e-7)
}
+128
View File
@@ -0,0 +1,128 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Concat joins u after t along an existing axis, the differentiable
// inverse of Slice, and the building block that lets recurrent layers
// assemble per-step outputs into one sequence core.
// Concat returns the tensors joined along the given existing dimension.
// Every other dimension must agree. The backward routes each side its
// own span of the incoming gradient along that dimension, narrowed to
// the side's dtype first: a real side of a real/complex join receives
// 2·Re(g), the rule every other mixed-dtype op applies.
func (t *Tensor) Concat(u *Tensor, dim int) (*Tensor, error) {
if err := t.checkDiff("Concat"); err != nil {
return nil, err
}
if err := u.checkDiff("Concat"); err != nil {
return nil, err
}
out, err := core.Concat(t.data, u.data, dim)
if err != nil {
return nil, err
}
at, au := t.data, u.data
return binaryResult("Concat", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
dt := at.Dtype()
du := au.Dtype()
// One narrowing per side before either span is copied: the
// concat's own dtype promotes along the ladder, so a complex
// gradient reaches a real operand in a mixed join and must
// narrow by 2·Re exactly as the leaf commit would. The narrowed
// arrays also put both copies on the dtype-matched raw path.
gA, err := narrowGradient(g, dt)
if err != nil {
return err
}
gB, err := narrowGradient(g, du)
if err != nil {
return err
}
shA, shB := at.Shape(), au.Shape()
outer, inner, err := outerInner(shA, dim)
if err != nil {
return err
}
spanA := shA[dim]
total := spanA + shB[dim]
da := gradSlot{arr: ar.borrowGrad(dt, shA), sh: shA}
// Every row contributes one contiguous inner run, so matching
// dtypes collapse the triple loop to a raw slice move per row.
fastA := !gA.arr.Strided() && gA.arr.Dtype() == dt && dt != core.Int
for o := range outer {
for i := range spanA {
d := o*spanA*inner + i*inner
s := o*total*inner + i*inner
if fastA {
copySegRaw(da.arr, gA.arr, d, s, inner)
continue
}
for j := range inner {
copyElem(da.arr, d+j, gA.arr, s+j)
}
}
}
db := gradSlot{arr: ar.borrowGrad(du, shB), sh: shB}
spanB := shB[dim]
fastB := !gB.arr.Strided() && gB.arr.Dtype() == du && du != core.Int
for o := range outer {
for i := range spanB {
d := o*spanB*inner + i*inner
s := o*total*inner + (spanA+i)*inner
if fastB {
copySegRaw(db.arr, gB.arr, d, s, inner)
continue
}
for j := range inner {
copyElem(db.arr, d+j, gB.arr, s+j)
}
}
}
dst[0], dst[1] = da, db
return nil
}), nil
}
// copySegRaw moves n elements from src at sOff to dst at dOff through
// the raw payloads. The caller checks dtype equality and contiguity;
// the per-element values, and so the bits, are the ones copyElem
// writes one accessor call at a time.
func copySegRaw(dst, src *core.Array, dOff, sOff, n int) {
switch src.Dtype() {
case core.Float32:
copy(dst.RawFloat32s()[dOff:dOff+n], src.RawFloat32s()[sOff:sOff+n])
case core.Float:
copy(dst.RawFloats()[dOff:dOff+n], src.RawFloats()[sOff:sOff+n])
case core.Complex:
copy(dst.RawComplexes()[dOff:dOff+n], src.RawComplexes()[sOff:sOff+n])
default:
copy(dst.RawInts()[dOff:dOff+n], src.RawInts()[sOff:sOff+n])
}
}
// outerInner splits a shape into the products of the dimensions before
// and after dim, the strides a flat row-major walk needs when only one
// axis is being split or joined.
func outerInner(shape []int, dim int) (int, int, error) {
if len(shape) == 0 {
return 0, 0, errf("Concat: cannot concatenate a scalar")
}
if dim < 0 || dim >= len(shape) {
return 0, 0, errf("Concat: dimension %d is out of range for shape %v", dim, shape)
}
outer := 1
for d := range dim {
outer *= shape[d]
}
inner := 1
for d := dim + 1; d < len(shape); d++ {
inner *= shape[d]
}
return outer, inner, nil
}
+167
View File
@@ -0,0 +1,167 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
func TestTensorConcatForward(t *testing.T) {
a, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
b, _ := core.FromFloats([]float64{5, 6, 7, 8, 9, 10}, 2, 3)
out, err := FromArray(a, false).Concat(FromArray(b, false), 1)
if err != nil {
t.Fatal(err)
}
if got := out.Data().Shape(); got[0] != 2 || got[1] != 5 {
t.Fatalf("concat shape: %v", got)
}
want := []float64{1, 2, 5, 6, 7, 3, 4, 8, 9, 10}
for i := range want {
if g := out.Data().FloatAt(i); g != want[i] {
t.Fatalf("concat[%d] = %v, want %v", i, g, want[i])
}
}
// Concatenation along the leading axis stacks the blocks.
c, _ := core.FromFloats([]float64{1, 2}, 1, 2)
d, _ := core.FromFloats([]float64{3, 4}, 1, 2)
vert, err := FromArray(c, false).Concat(FromArray(d, false), 0)
if err != nil {
t.Fatal(err)
}
if vert.Data().Shape()[0] != 2 {
t.Fatalf("vertical shape: %v", vert.Data().Shape())
}
// Mismatched ranks and out-of-range axes error.
misrank, _ := core.Reshape(c, 2)
if _, err := FromArray(a, false).Concat(FromArray(misrank, false), 0); err == nil {
t.Fatal("rank mismatch accepted")
}
if _, err := FromArray(a, false).Concat(FromArray(b, false), 2); err == nil {
t.Fatal("out-of-range dimension accepted")
}
}
// TestTensorConcatGradients checks both backward spans against central
// differences with a weighted loss so every slot gets a distinct weight.
func TestTensorConcatGradients(t *testing.T) {
cases := []struct {
dim int
aVal, bVal []float64
aShape, bShape []int
waVal, wbVal []float64
}{
{
dim: 1,
aVal: []float64{0.5, -1, 2, 0.25}, aShape: []int{2, 2},
bVal: []float64{1.5, -0.5, 1, 2, -2, 0.75}, bShape: []int{2, 3},
waVal: []float64{0.1, -0.4, 0.9, 0.6},
wbVal: []float64{0.2, 0.3, -0.7, 0.8, 0.05, -0.6},
},
{
dim: 0,
aVal: []float64{0.3, 1, -0.25, 2}, aShape: []int{2, 2},
bVal: []float64{-1.5, 0.4, 0.9, 1, -2, 0.7}, bShape: []int{3, 2},
waVal: []float64{0.55, -0.35, 0.85, 0.15},
wbVal: []float64{0.45, -0.65, 0.95, 0.05, -0.5, 0.75},
},
}
for _, tc := range cases {
a, _ := core.FromFloats(tc.aVal, tc.aShape...)
b, _ := core.FromFloats(tc.bVal, tc.bShape...)
wa, _ := core.FromFloats(tc.waVal, tc.aShape...)
wb, _ := core.FromFloats(tc.wbVal, tc.bShape...)
at := FromArray(a, true)
bt := FromArray(b, true)
joint, err := at.Concat(bt, tc.dim)
if err != nil {
t.Fatal(err)
}
wc, _ := core.Concat(wa, wb, tc.dim)
scaled, err := joint.Mul(FromArray(wc, false))
if err != nil {
t.Fatal(err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
fa := func(v *core.Array) float64 { return weightedConcatSum(v, b, wa, wb, tc.dim) }
fb := func(v *core.Array) float64 { return weightedConcatSum(a, v, wa, wb, tc.dim) }
checkSpan(t, at.Grad(), numericGrad(fa, a))
checkSpan(t, bt.Grad(), numericGrad(fb, b))
}
}
// TestTensorConcatGradientDtype keeps each side's gradient in its own
// element type: float32 inputs never come back as float64 leaves.
func TestTensorConcatGradientDtype(t *testing.T) {
gen := core.NewGenerator(5)
af, _ := core.Float32s(gen, 6)
afArr, _ := core.Reshape(af, 2, 3)
bf, _ := core.Float32s(gen, 6)
bfArr, _ := core.Reshape(bf, 2, 3)
at := FromArray(afArr, true)
bt := FromArray(bfArr, true)
joint, err := at.Concat(bt, 1)
if err != nil {
t.Fatal(err)
}
loss, err := joint.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
if at.Grad().Dtype() != core.Float32 || bt.Grad().Dtype() != core.Float32 {
t.Fatalf("gradient dtypes: %v and %v", at.Grad().Dtype(), bt.Grad().Dtype())
}
for i := range at.Grad().Len() {
if at.Grad().FloatAt(i) != 1 {
t.Errorf("float32 gradient slot %d: %v, want 1", i, at.Grad().FloatAt(i))
}
}
}
// weightedConcatSum evaluates Σ w∘Concat(x, y, dim) with fixed weights,
// the scalar objective whose gradients the backward is checked against.
func weightedConcatSum(x, y *core.Array, wx, wy *core.Array, dim int) float64 {
joint, err := FromArray(x, false).Concat(FromArray(y, false), dim)
if err != nil {
return math.NaN()
}
wc, _ := core.Concat(wx, wy, dim)
total := 0.0
for i := range joint.Data().Len() {
total += wc.FloatAt(i) * joint.Data().FloatAt(i)
}
return total
}
// checkSpan reports every slot where the analytic gradient drifts from
// the central-difference reference.
func checkSpan(t *testing.T, got *core.Array, ref []float64) {
t.Helper()
if got.Len() != len(ref) {
t.Fatalf("gradient length %d, reference %d", got.Len(), len(ref))
}
for i := range ref {
if math.Abs(got.FloatAt(i)-ref[i]) > 1e-5 {
t.Errorf("gradient[%d] = %v, want ≈%v", i, got.FloatAt(i), ref[i])
}
}
}
+200
View File
@@ -0,0 +1,200 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
func TestTensorSqueezeUnsqueezeClip(t *testing.T) {
// Squeeze/Unsqueeze round-trip with gradient.
x, _ := core.FromFloats([]float64{1, 2, 3, 4}, 1, 4, 1)
xt := FromArray(x, true)
sq, err := xt.Squeeze(2)
if err != nil {
t.Fatal(err)
}
if sq.Data().NDim() != 2 {
t.Fatalf("Squeeze ndim: %d", sq.Data().NDim())
}
back, err := sq.Unsqueeze(2)
if err != nil {
t.Fatal(err)
}
s, _ := back.Sum()
if err := s.Backward(); err != nil {
t.Fatal(err)
}
for i := range 4 {
if g := xt.Grad().FloatAt(i); g != 1 {
t.Errorf("Squeeze/Unsqueeze grad[%d]: %v, want 1", i, g)
}
}
// Clip gradient: 1 inside [lo, hi], 0 outside.
c, _ := core.FromFloats([]float64{-1, 0.5, 2}, 3)
ct := FromArray(c, true)
cl, err := ct.Clip(0, 1)
if err != nil {
t.Fatal(err)
}
s2, _ := cl.Sum()
if err := s2.Backward(); err != nil {
t.Fatal(err)
}
want := []float64{0, 1, 0}
for i := range 3 {
if g := ct.Grad().FloatAt(i); g != want[i] {
t.Errorf("Clip grad[%d]: %v, want %v", i, g, want[i])
}
}
}
func TestAxisReductionAutograd(t *testing.T) {
x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
xt := FromArray(x, true)
s, err := xt.SumAxis(1)
if err != nil {
t.Fatal(err)
}
if s.Data().Len() != 2 {
t.Fatalf("SumAxis len: %d", s.Data().Len())
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
for i := range x.Len() {
if g := xt.Grad().FloatAt(i); g != 1 {
t.Errorf("SumAxis grad[%d]: %v, want 1", i, g)
}
}
mean, err := FromArray(x, true).MeanAxis(1)
if err != nil {
t.Fatal(err)
}
if err := mean.Backward(); err != nil {
t.Fatal(err)
}
}
func TestL2NormAxisAutogradGradient(t *testing.T) {
xv := []float64{3, 4, 0.5, 0.5}
x, _ := core.FromFloats(xv, 1, 1, 2, 2)
xt := FromArray(x, true)
out, err := xt.L2NormAxis(1)
if err != nil {
t.Fatal(err)
}
s, _ := out.Sum()
if err := s.Backward(); err != nil {
t.Fatal(err)
}
analytic := make([]float64, x.Len())
for i := range x.Len() {
analytic[i] = xt.Grad().FloatAt(i)
}
ref := numericGrad(func(a *core.Array) float64 {
o, err := FromArray(a, false).L2NormAxis(1)
if err != nil {
t.Fatal(err)
}
ss, _ := o.Sum()
return ss.Data().FloatAt(0)
}, x)
if d := maxAbsDiff(analytic, ref); d > 1e-6 {
t.Errorf("L2NormAxis grad: max diff %v", d)
}
}
func TestBroadcastToAutograd(t *testing.T) {
x, _ := core.FromFloats([]float64{1, 2, 3}, 1, 3)
xt := FromArray(x, true)
out, err := xt.BroadcastTo(2, 3)
if err != nil {
t.Fatal(err)
}
if out.Data().Shape()[0] != 2 {
t.Fatalf("BroadcastTo shape: %v", out.Data().Shape())
}
onesArr, _ := core.Ones(core.Float, 2, 3)
loss, _ := out.Mul(FromArray(onesArr, false))
s, _ := loss.Sum()
if err := s.Backward(); err != nil {
t.Fatal(err)
}
// Gradient sums over replicated rows.
for i := range 3 {
if g := xt.Grad().FloatAt(i); g != 2 {
t.Errorf("BroadcastTo grad[%d]: %v, want 2", i, g)
}
}
}
func TestPowAbsSqrtFloorAutogradGradient(t *testing.T) {
cases := []struct {
name string
vals []float64
fn func(*Tensor) (*Tensor, error)
}{
{"Pow3", []float64{0.5, 1.5}, func(x *Tensor) (*Tensor, error) { return x.Pow(3) }},
{"Abs", []float64{0.5, -1.5}, func(x *Tensor) (*Tensor, error) { return x.Abs() }},
{"Sqrt", []float64{0.25, 2.25}, func(x *Tensor) (*Tensor, error) { return x.Sqrt() }},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
a, _ := core.FromFloats(tc.vals, 2)
at := FromArray(a, true)
out, err := tc.fn(at)
if err != nil {
t.Fatal(err)
}
s, err := out.Sum()
if err != nil {
t.Fatal(err)
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
analytic := make([]float64, a.Len())
for i := range a.Len() {
analytic[i] = at.Grad().FloatAt(i)
}
ref := numericGrad(func(v *core.Array) float64 {
o, err := tc.fn(FromArray(v, false))
if err != nil {
t.Fatal(err)
}
ss, err := o.Sum()
if err != nil {
t.Fatal(err)
}
return ss.Data().FloatAt(0)
}, a)
if d := maxAbsDiff(analytic, ref); d > 1e-6 {
t.Errorf("%s grad: max diff %v", tc.name, d)
}
})
}
// Floor contributes no gradient.
a, _ := core.FromFloats([]float64{1.4, 2.6}, 2)
at := FromArray(a, true)
fl, err := at.Floor()
if err != nil {
t.Fatal(err)
}
s, _ := fl.Sum()
if err := s.Backward(); err != nil {
t.Fatal(err)
}
for i := range a.Len() {
if g := at.Grad().FloatAt(i); g != 0 {
t.Errorf("Floor grad[%d]: %v, want 0", i, g)
}
}
}
+82
View File
@@ -0,0 +1,82 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package grad is reverse-mode automatic differentiation over the array
// surface: a computation written as ordinary Go calls over [Tensor]
// values records a graph, and one call to [Tensor.Backward] propagates
// the gradient from the output back to every leaf that requires it.
//
// # The graph
//
// Leaves come from [FromFloat64s] or [FromArray]. Every differentiable
// method records the operation it performs together with its inputs, and
// Backward sweeps the recorded nodes in reverse, applying each node's
// adjoint. The method set covers the arithmetic, the matrix products
// (single and batched), the element-wise transcendentals, the
// reductions, slicing, concatenation, axis permutation and the Fourier
// transforms, with [Tensor.Conj], [Tensor.Real], [Tensor.Imag] and
// [Tensor.Abs2] carrying complex values into a real loss.
//
// loss, err := total.Mean() // the forward pass records the nodes
// if err != nil {
// return err
// }
// if err := loss.Backward(); err != nil {
// return err
// }
// x.Grad() // the accumulated gradient
//
// # The contract
//
// - float, float32 and complex128 tensors differentiate; an int tensor
// is refused by every differentiable method.
// - The loss must be real: Backward seeds the output with ones, and a
// complex output is rejected with an error naming Real, Imag, Abs
// and Abs2 as the reducers that turn it into a real scalar.
// - Gradients carry the leaf's dtype. A mixed-dtype graph narrows each
// gradient to the dtype of the tensor it accumulates into before the
// leaf is written.
// - Backward accumulates into the gradients already present, so
// [Tensor.ZeroGrad] precedes a fresh pass unless accumulation is
// wanted.
// - The graph is rebuilt on every forward pass. Each operation's
// backward closure captures the operands as they were when the
// operation ran, so a [Tensor.ReplaceWith] afterwards changes the
// next pass and not the recorded one.
//
// # Complex graphs
//
// Complex tensors differentiate under the Wirtinger convention: the
// gradient a complex leaf accumulates is ∂L/∂z̄, the coefficient g of
// dL = 2·Re(g·dz), which is the direction gradient descent steps along.
// The adjoint of a holomorphic y = f(z) is therefore dz = g·conj(f′(z)),
// and every complex adjoint in this package conjugates exactly where the
// calculus puts it. A real tensor inside a complex graph narrows the
// incoming gradient by 2·Re, the factor that also cancels the ½ the Real
// and Imag backward paths contribute, so a graph mixing the two dtypes
// composes exactly.
//
// # Second-order and solver tools
//
// On top of the graph sit the second derivative and the methods that
// need one: [Hessian] (dense, 2n gradient evaluations),
// [HessianVectorProduct] (H·v in two gradient evaluations),
// [MinimiseNewtonCG] (truncated conjugate gradients on the Hessian
// system with an Armijo line search), [SampleHMC] (Hamiltonian Monte
// Carlo on any differentiable unnormalised log density) and [AdjointODE]
// (adjoint sensitivities of an ODE solution at the cost of one extra
// solve).
//
// Each of them differentiates through a reverse pass that commits
// nothing, so the accumulated gradients of the tensors the caller's
// closure holds are left exactly as they were, whether the call succeeds
// or fails.
//
// # What it does not do
//
// There is no forward-mode differentiation, nothing beyond the second
// derivative, no graph serialisation and no parameter registry: the
// caller owns the leaves and the graph is a transient record of one
// forward pass. The operation set is closed, so a new primitive is a new
// method here and never a user-registered op.
package grad
+218
View File
@@ -0,0 +1,218 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad_test
// The godoc examples for the autograd package: the flagship workflows
// as runnable, checked snippets. Each one pins the numbers it prints,
// so a change in the adjoint of an op or in a solver's default shows
// up as a failing example rather than as stale prose.
import (
"fmt"
"log"
tensor "sourcedock.dev/petrbalvin/tensor"
"sourcedock.dev/petrbalvin/tensor/grad"
)
// A scalar loss by hand on a small graph: z = Σ x² over a two-element
// leaf, differentiated in one reverse sweep. The leaf is entered twice
// by the product, so the backward adds both contributions and the
// answer is the 2x the calculus gives.
func ExampleTensor_Backward() {
x, err := grad.FromFloat64s([]float64{2, 3}, true, 2)
if err != nil {
log.Fatal(err)
}
sq, err := x.Mul(x)
if err != nil {
log.Fatal(err)
}
loss, err := sq.Sum()
if err != nil {
log.Fatal(err)
}
if err := loss.Backward(); err != nil {
log.Fatal(err)
}
fmt.Println(loss.Data(), x.Grad())
// Output: float (1) [13] float (2) [4, 6]
}
// A matrix product and a reduction as graph nodes: the gradient of the
// sum of A·B is ones·Bᵀ, one row sum of B per row of A.
func ExampleTensor_MatMul() {
a, err := grad.FromFloat64s([]float64{1, 2, 3, 4, 5, 6}, true, 2, 3)
if err != nil {
log.Fatal(err)
}
b, err := grad.FromFloat64s([]float64{1, 0, 0, 1, 1, 1}, false, 3, 2)
if err != nil {
log.Fatal(err)
}
prod, err := a.MatMul(b)
if err != nil {
log.Fatal(err)
}
total, err := prod.Sum()
if err != nil {
log.Fatal(err)
}
if err := total.Backward(); err != nil {
log.Fatal(err)
}
fmt.Println(total.Data(), a.Grad())
// Output: float (1) [30] float (2, 3) [1, 1, 2, 1, 1, 2]
}
// A complex graph with a real loss. The leaf is complex, the loss is
// Σ|z|², and the gradient a complex leaf accumulates is ∂L/∂z̄, which
// for |z|² is z itself.
func ExampleTensor_Abs2() {
data, err := tensor.FromComplexes([]complex128{1 + 2i, 3 - 1i}, 2)
if err != nil {
log.Fatal(err)
}
z := grad.FromArray(data, true)
magnitude, err := z.Abs2()
if err != nil {
log.Fatal(err)
}
loss, err := magnitude.Sum()
if err != nil {
log.Fatal(err)
}
if err := loss.Backward(); err != nil {
log.Fatal(err)
}
fmt.Println(loss.Data(), z.Grad())
// Output: float (1) [15] complex (2) [(1+2i), (3-1i)]
}
// The dense second derivative of Σ x² at (1, 2): the Hessian of a
// quadratic form is twice its matrix, here 2·I.
func ExampleHessian() {
f := func(x *grad.Tensor) (*grad.Tensor, error) {
sq, err := x.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
point, err := grad.FromFloat64s([]float64{1, 2}, true, 2)
if err != nil {
log.Fatal(err)
}
h, err := grad.Hessian(f, point, grad.HessianOptions{})
if err != nil {
log.Fatal(err)
}
fmt.Println(h)
// Output: float (2, 2) [2, 0, 0, 2]
}
// The same function and point contracted with a direction: H·v in two
// gradient evaluations instead of the dense Hessian's four.
func ExampleHessianVectorProduct() {
f := func(x *grad.Tensor) (*grad.Tensor, error) {
sq, err := x.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
point, err := grad.FromFloat64s([]float64{1, 2}, true, 2)
if err != nil {
log.Fatal(err)
}
direction, err := grad.FromFloat64s([]float64{1, 1}, false, 2)
if err != nil {
log.Fatal(err)
}
hv, err := grad.HessianVectorProduct(f, point, direction, grad.HessianOptions{})
if err != nil {
log.Fatal(err)
}
// A central difference along the direction, so the answer carries
// the rounding of the two gradient evaluations it is built from.
fmt.Printf("(%.4f, %.4f)\n", hv.FloatAt(0), hv.FloatAt(1))
// Output: (2.0000, 2.0000)
}
// Newton-CG on the quadratic Σ (x − c)², whose minimiser is c and
// whose value there is zero. The curvature comes from the
// Hessian-vector product, so no dense Hessian is ever formed.
func ExampleMinimiseNewtonCG() {
centre, err := grad.FromFloat64s([]float64{1.5, -2.5}, false, 2)
if err != nil {
log.Fatal(err)
}
f := func(x *grad.Tensor) (*grad.Tensor, error) {
diff, err := x.Sub(centre)
if err != nil {
return nil, err
}
sq, err := diff.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
x0, err := tensor.FromFloats([]float64{0, 0}, 2)
if err != nil {
log.Fatal(err)
}
x, value, err := grad.MinimiseNewtonCG(f, x0, grad.NewtonCGOptions{})
if err != nil {
log.Fatal(err)
}
fmt.Printf("x = (%.4f, %.4f), f = %.4f\n", x.FloatAt(0), x.FloatAt(1), value)
// Output: x = (1.5000, -2.5000), f = 0.0000
}
// Hamiltonian Monte Carlo on the two-dimensional standard normal,
// whose log density is −‖q‖²/2. The seed makes the chain reproducible,
// so the moments of the first component are fixed numbers and not a
// range: the target has mean zero and variance one.
func ExampleSampleHMC() {
logDensity := func(q *grad.Tensor) (*grad.Tensor, error) {
sq, err := q.Abs2()
if err != nil {
return nil, err
}
total, err := sq.Sum()
if err != nil {
return nil, err
}
return total.Scale(-0.5)
}
q0, err := tensor.FromFloats([]float64{2, -2}, 2)
if err != nil {
log.Fatal(err)
}
samples, err := grad.SampleHMC(logDensity, q0, grad.HMCOptions{
Step: 0.25,
Steps: 16,
BurnIn: 500,
Thin: 1,
Samples: 2000,
Seed: 7,
})
if err != nil {
log.Fatal(err)
}
rows := samples.Shape()[0]
mean := 0.0
for row := range rows {
mean += samples.FloatAt(row * 2)
}
mean /= float64(rows)
variance := 0.0
for row := range rows {
d := samples.FloatAt(row*2) - mean
variance += d * d / float64(rows)
}
fmt.Printf("%v: mean %.3f, variance %.3f\n", samples.Shape(), mean, variance)
// Output: [2000 2]: mean -0.012, variance 1.013
}
+187
View File
@@ -0,0 +1,187 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Property fuzz targets over index-moving operations. `go test` runs
// the seed corpus on every commit; longer campaigns run under
// -fuzz=Fuzz<Name> when a shape-handling change lands.
// FuzzTransposeAxesRoundTrip drives random permutations through the
// axis move and its inverse: whatever valid permutation arrives, the
// double transpose must restore the exact element order.
func FuzzTransposeAxesRoundTrip(f *testing.F) {
f.Add([]byte{0, 1}, 6)
f.Add([]byte{1, 0}, 6)
f.Add([]byte{2, 0, 1}, 8)
f.Fuzz(func(t *testing.T, permBytes []byte, total int) {
if total <= 0 || total > 4096 {
t.Skip()
}
rank := len(permBytes)
switch rank {
case 2:
total -= total % 2
case 3:
total -= total % 4
default:
t.Skip()
}
if total == 0 {
t.Skip()
}
vals := make([]float64, total)
for i := range vals {
vals[i] = float64(i)
}
var shape []int
if rank == 2 {
shape = []int{total / 2, 2}
} else {
shape = []int{total / 4, 2, 2}
}
a, _ := core.FromFloats(vals, shape...)
xt := FromArray(a, false)
dims := make([]int, rank)
for i, pb := range permBytes {
dims[i] = int(pb) % rank
}
moved, err := xt.TransposeAxes(dims...)
if err != nil {
return // duplicate axes rejected by validation, fine
}
back, err := moved.TransposeAxes(inversePerm(dims)...)
if err != nil {
t.Fatalf("inverse of %v failed: %v", dims, err)
}
for i := range vals {
if back.Data().FloatAt(i) != vals[i] {
t.Fatalf("round trip lost element %d", i)
}
}
})
}
// FuzzOneHotContracts checks both sides of the encoder contract for
// arbitrary code sets: in-range codes yield exactly one hot cell per
// row, any out-of-range code is a loud error. One input byte splits
// into a high bit forcing negativity plus a low-bit class selector.
func FuzzOneHotContracts(f *testing.F) {
f.Add([]byte{0, 1, 2}, uint8(3))
f.Add([]byte{5}, uint8(8))
f.Add([]byte{200, 201}, uint8(3))
f.Fuzz(func(t *testing.T, raw []byte, classByte uint8) {
classes := int(classByte)%9 + 1
codes := make([]int64, len(raw))
valid := true
for i, b := range raw {
c := int64(b)
if b >= 128 { // force some negative probes
c = -int64(b - 127)
} else {
c %= int64(classes)
}
if c < 0 || c >= int64(classes) {
valid = false
}
codes[i] = c
}
arr, _ := core.FromInts(codes, len(codes))
hot, err := core.OneHot(arr, classes)
if !valid {
if err == nil {
t.Fatalf("invalid codes accepted for %d classes", classes)
}
return
}
if err != nil {
t.Fatalf("valid codes rejected: %v", err)
}
for i := range len(codes) {
sum := 0.0
for j := range classes {
sum += float64(hot.FloatAt(i*classes + j))
}
if sum != 1 {
t.Fatalf("row %d sums to %v, want one hot cell", i, sum)
}
}
})
}
// FuzzConcatSplitGradientConserves mass: splitting a concatenated
// output's gradient must hand every element back to its own side with
// coefficient exactly one, for whatever layout the corpus invents.
func FuzzConcatSplitGradientConserves(f *testing.F) {
f.Add([]byte{1, 2, 3, 4}, uint8(2))
f.Add([]byte{9, 7, 5}, uint8(1))
f.Add([]byte{10, 20, 30, 40, 50, 60, 70, 80}, uint8(0))
f.Add([]byte{11, 21, 31, 41, 51, 61, 71, 81, 91, 101}, uint8(4))
f.Fuzz(func(t *testing.T, raw []byte, rowsByte uint8) {
rows := int(rowsByte)%7 + 1 // left matrix rows, 1..7
leftLen := rows * 2
if len(raw) <= leftLen {
t.Skip()
}
rightRows := (len(raw) - leftLen) / 2
leftVals := make([]float64, leftLen)
for i := range leftVals {
leftVals[i] = float64(raw[i])
}
rightVals := make([]float64, rightRows*2)
for i := range rightVals {
rightVals[i] = float64(raw[leftLen+i])
}
a, err := core.FromFloats(leftVals, rows, 2)
if err != nil {
t.Skip()
}
b, err := core.FromFloats(rightVals, rightRows, 2)
if err != nil {
t.Skip()
}
at := FromArray(a, true)
bt := FromArray(b, true)
joint, cerr := at.Concat(bt, 0)
if cerr != nil {
t.Fatal(cerr)
}
loss, serr := joint.Sum()
if serr != nil {
t.Fatal(serr)
}
if berr := loss.Backward(); berr != nil {
t.Fatal(berr)
}
ga, gb := at.Grad(), bt.Grad()
if ga.Len() != a.Len() || gb.Len() != b.Len() {
t.Fatalf("gradient spans drifted: %d+%d vs %d+%d",
ga.Len(), gb.Len(), a.Len(), b.Len())
}
for i := range ga.Len() {
if ga.FloatAt(i) != 1 {
t.Fatalf("left span slot %d = %v", i, ga.FloatAt(i))
}
}
for i := range gb.Len() {
if gb.FloatAt(i) != 1 {
t.Fatalf("right span slot %d = %v", i, gb.FloatAt(i))
}
}
})
}
+875
View File
@@ -0,0 +1,875 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression pins for the grad package guards: refusals and gradients
// that used to panic or silently pass through. Each test names
// the defect it pins and fails without its fix.
// TestConcatMixedDtypeBackwardNarrows pins the mixed real/complex
// Concat backward. The join promotes along the dtype ladder, so a
// complex gradient reaches a real operand; that operand must receive
// 2·Re of its span, the package's real-operand rule, instead of dying
// in copyElem, which used to read the complex source through FloatAt (a
// nil int payload) and panic. Both operand orders and the constant
// complex operand are covered, and the real side is checked against
// central differences as well as its closed form.
func TestConcatMixedDtypeBackwardNarrows(t *testing.T) {
xv := []float64{1.5, -0.75, 2.25, 0.5, -1, 3}
zv := []complex128{1 + 1i, 2 - 0.5i, -0.25 + 0.75i, 3, 0.5 - 2i, -1.5i}
// L = Σ|Concat(a, b)|² = Σx² + Σ|z|², so the real operand's
// gradient is 2x and the complex one's is z, whatever the axis.
cases := []struct {
name string
dim int
xSh []int
zSh []int
}{
{name: "dim0", dim: 0, xSh: []int{2, 3}, zSh: []int{2, 3}},
{name: "dim1", dim: 1, xSh: []int{3, 2}, zSh: []int{3, 2}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
x, err := core.FromFloats(xv, tc.xSh...)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
z, err := core.FromComplexes(zv, tc.zSh...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
// The A side of a real x complex join is the panic the
// report reproduced; the B side of a complex x real join
// fails identically.
concatChecked(t, "real||cplx", FromArray(x, true), FromArray(z, true), tc.dim, xv, zv)
concatChecked(t, "cplx||real", FromArray(z, true), FromArray(x, true), tc.dim, xv, zv)
// A constant complex operand takes the same span loop, so
// the real side must still narrow.
concatChecked(t, "real||cplxconst", FromArray(x, true), FromArray(z, false), tc.dim, xv, nil)
// Finite differences confirm the 2·Re rule itself, not just
// its agreement with the closed form.
ref := numericGrad(func(v *core.Array) float64 {
realSide := FromArray(v, false)
cat, cerr := realSide.Concat(FromArray(z, false), tc.dim)
if cerr != nil {
return math.NaN()
}
sq, cerr := cat.Abs2()
if cerr != nil {
return math.NaN()
}
s, cerr := sq.Sum()
if cerr != nil {
return math.NaN()
}
return s.Data().FloatAt(0)
}, x)
xt := FromArray(x, true)
concatChecked(t, "fd", xt, FromArray(z, false), tc.dim, xv, nil)
if d := maxAbsDiff(flatFloats(xt.Grad()), ref); d > 1e-8 {
t.Errorf("real operand gradient differs from central differences by %g", d)
}
})
}
// A float32 real operand beside a complex one: the narrowed side
// keeps the operand's width, and the value is still 2x.
t.Run("float32RealSide", func(t *testing.T) {
x32 := []float32{1.5, -0.75, 2.25, 0.5}
z32 := []complex128{1 + 1i, 2 - 0.5i, -0.25 + 0.75i, 3}
a, err := core.FromFloat32s(x32, 2, 2)
if err != nil {
t.Fatalf("FromFloat32s: %v", err)
}
z, err := core.FromComplexes(z32, 2, 2)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
xt := FromArray(a, true)
cat, err := xt.Concat(FromArray(z, true), 0)
if err != nil {
t.Fatalf("Concat: %v", err)
}
sq, err := cat.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if got := g.Dtype(); got != core.Float32 {
t.Errorf("float32 real operand gradient dtype = %s, want float32", got)
}
for i, v := range x32 {
if got, want := g.FloatAt(i), 2*float64(v); math.Abs(got-want) > 1e-6 {
t.Errorf("float32 gradient[%d] = %v, want %v (2·Re rule)", i, got, want)
}
}
})
// A rebased-view complex operand: the view aliases a longer payload,
// and the gradient must scatter back through the view's own slots.
t.Run("rebasedViewOperand", func(t *testing.T) {
big, err := core.FromComplexes(zv[:4], 4)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
zb := FromArray(big, true)
view, err := zb.Slice(0, 1, 3)
if err != nil {
t.Fatalf("Slice: %v", err)
}
x, err := core.FromFloats(xv[:2], 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
cat, err := xt.Concat(view, 0)
if err != nil {
t.Fatalf("Concat: %v", err)
}
sq, err := cat.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
for i := range xv[:2] {
if got, want := xt.Grad().FloatAt(i), 2*xv[i]; got != want {
t.Errorf("view case: real gradient[%d] = %v, want %v", i, got, want)
}
}
g := zb.Grad()
if g == nil {
t.Fatal("the view's parent received no gradient")
}
// Only the aliased slots carry the view's own values; the rest
// stays zero through the Slice backward.
want := []complex128{0, zv[1], zv[2], 0}
for i := range want {
if got := g.ComplexAt(i); got != want[i] {
t.Errorf("view case: parent gradient[%d] = %v, want %v", i, got, want[i])
}
}
})
// The mirror image: a rebased-view REAL operand beside a complex one.
// The narrowed span (2·Re) reaches the view as float and Slice's
// backward scatters it into the parent's own slots, leaving the rest
// zero. This is the one shape that exercises both fixes together.
t.Run("rebasedViewRealOperand", func(t *testing.T) {
full := []float64{9, 1.5, -0.75, 7}
fa, err := core.FromFloats(full, 4)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
fb := FromArray(fa, true)
view, err := fb.Slice(0, 1, 3)
if err != nil {
t.Fatalf("Slice: %v", err)
}
z, err := core.FromComplexes(zv[:2], 2)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
cat, err := view.Concat(FromArray(z, true), 0)
if err != nil {
t.Fatalf("Concat: %v", err)
}
sq, err := cat.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := fb.Grad()
if g == nil {
t.Fatal("the view's parent received no gradient")
}
want := []float64{0, 2 * full[1], 2 * full[2], 0}
for i := range want {
if got := g.FloatAt(i); got != want[i] {
t.Errorf("view case: parent gradient[%d] = %v, want %v", i, got, want[i])
}
}
})
}
// concatChecked builds Σ|first.Concat(second, dim)|², runs the
// backward and checks every gradient-carrying operand against its
// closed form by dtype: a real side against 2x (the 2·Re rule), a
// complex side against z. A nil wantZ marks a constant complex operand,
// and a non-grad operand has nothing to check.
func concatChecked(t *testing.T, label string, first, second *Tensor, dim int, xv []float64, wantZ []complex128) {
t.Helper()
cat, err := first.Concat(second, dim)
if err != nil {
t.Fatalf("%s: Concat: %v", label, err)
}
sq, err := cat.Abs2()
if err != nil {
t.Fatalf("%s: Abs2: %v", label, err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("%s: Sum: %v", label, err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("%s: Backward: %v", label, err)
}
for _, side := range []*Tensor{first, second} {
if !side.RequiresGrad() {
continue
}
g := side.Grad()
if g == nil {
t.Fatalf("%s: an operand received no gradient", label)
}
if g.Dtype() == core.Complex {
if wantZ == nil {
continue
}
for i := range wantZ {
if got := g.ComplexAt(i); got != wantZ[i] {
t.Errorf("%s: complex operand gradient[%d] = %v, want %v", label, i, got, wantZ[i])
}
}
continue
}
for i := range xv {
if got, want := g.FloatAt(i), 2*xv[i]; got != want {
t.Errorf("%s: real operand gradient[%d] = %v, want %v (the 2·Re rule)", label, i, got, want)
}
}
}
}
// TestHessianAndHVPRejectComplexOperands pins the dtype guard on the
// two second-order helpers: a complex point or direction is refused
// with an error naming the dtype, as MinimiseNewtonCG, SampleHMC and
// AdjointODE already do, never a panic out of flatFloats.
func TestHessianAndHVPRejectComplexOperands(t *testing.T) {
z, err := core.FromComplexes([]complex128{1 + 1i, 2 - 1i}, 2)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
zt := FromArray(z, false)
xf, err := core.FromFloats([]float64{1, 2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(xf, false)
objective := func(q *Tensor) (*Tensor, error) {
sq, err := q.Abs2()
if err != nil {
return nil, err
}
return sq.Sum()
}
refusesComplex(t, "Hessian on a complex point", func() error {
_, err := Hessian(objective, zt, HessianOptions{})
return err
})
refusesComplex(t, "HessianVectorProduct on a complex point", func() error {
_, err := HessianVectorProduct(objective, zt, xt, HessianOptions{})
return err
})
refusesComplex(t, "HessianVectorProduct on a complex direction", func() error {
_, err := HessianVectorProduct(objective, xt, zt, HessianOptions{})
return err
})
// The guards must not narrow the accepted surface: a real point
// still differentiates, and the real callers keep working.
h, err := Hessian(objective, xt, HessianOptions{})
if err != nil {
t.Fatalf("Hessian of a real point: %v", err)
}
// Σ|q|² over a real q is Σq², whose Hessian is 2·I.
for i := range 2 {
if got := h.FloatAt(i*2 + i); math.Abs(got-2) > 1e-6 {
t.Errorf("real Hessian diagonal[%d] = %v, want 2", i, got)
}
}
}
// refusesComplex runs fn and requires an error that names the
// complex dtype, treating a panic as the failure it reports.
func refusesComplex(t *testing.T, label string, fn func() error) {
t.Helper()
defer func() {
if r := recover(); r != nil {
t.Errorf("%s panicked: %v", label, r)
}
}()
err := fn()
if err == nil {
t.Errorf("%s was accepted", label)
return
}
if !strings.Contains(err.Error(), "complex") {
t.Errorf("%s error %q does not name the dtype", label, err)
}
}
// TestSecondOrderHelpersLeaveCallerGradients pins the internal reverse
// passes of Hessian, HessianVectorProduct and MinimiseNewtonCG: the
// objectives close over the caller's trainable tensors, and no path
// (success or error) may mutate their accumulated gradients. AdjointODE
// documents and implements the same guarantee.
func TestSecondOrderHelpersLeaveCallerGradients(t *testing.T) {
thetaVals := []float64{2, 3}
presetVals := []float64{0.5, 0.25}
thetaArr, err := core.FromFloats(thetaVals, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xArr, err := core.FromFloats([]float64{1, -1}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
// Σ z_i²·θ_i: the Hessian is diag(2θ) = diag(4, 6).
weighted := func(theta *Tensor) func(*Tensor) (*Tensor, error) {
return func(z *Tensor) (*Tensor, error) {
sq, err := z.Mul(z)
if err != nil {
return nil, err
}
w, err := sq.Mul(theta)
if err != nil {
return nil, err
}
return w.Sum()
}
}
// Σ θ_i²: independent of the probes, the disconnected-objective case.
constant := func(theta *Tensor) func(*Tensor) (*Tensor, error) {
return func(*Tensor) (*Tensor, error) {
sq, err := theta.Mul(theta)
if err != nil {
return nil, err
}
return sq.Sum()
}
}
t.Run("hessian success", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, presetBits := presetGrad(t, theta, presetVals)
h, err := Hessian(weighted(theta), FromArray(xArr, false), HessianOptions{})
if err != nil {
t.Fatalf("Hessian: %v", err)
}
for i := range 2 {
if got, want := h.FloatAt(i*2+i), 2*thetaVals[i]; math.Abs(got-want) > 1e-8 {
t.Errorf("Hessian diagonal[%d] = %v, want %v", i, got, want)
}
}
requireGradUntouched(t, "Hessian", theta, preset, presetBits)
})
t.Run("hessian error", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, presetBits := presetGrad(t, theta, presetVals)
if _, err := Hessian(constant(theta), FromArray(xArr, false), HessianOptions{}); err == nil {
t.Fatal("expected the disconnected-objective error")
}
requireGradUntouched(t, "Hessian error path", theta, preset, presetBits)
})
t.Run("hvp success", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, presetBits := presetGrad(t, theta, presetVals)
v, err := core.FromFloats([]float64{1, 0.5}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
hv, err := HessianVectorProduct(weighted(theta), FromArray(xArr, false), FromArray(v, false), HessianOptions{})
if err != nil {
t.Fatalf("HessianVectorProduct: %v", err)
}
for i := range 2 {
if got, want := hv.FloatAt(i), 2*thetaVals[i]*v.FloatAt(i); math.Abs(got-want) > 1e-6 {
t.Errorf("H·v[%d] = %v, want %v", i, got, want)
}
}
requireGradUntouched(t, "HessianVectorProduct", theta, preset, presetBits)
})
t.Run("hvp error", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, presetBits := presetGrad(t, theta, presetVals)
v, err := core.FromFloats([]float64{1, 0.5}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
if _, err := HessianVectorProduct(constant(theta), FromArray(xArr, false), FromArray(v, false), HessianOptions{}); err == nil {
t.Fatal("expected the disconnected-objective error")
}
requireGradUntouched(t, "HessianVectorProduct error path", theta, preset, presetBits)
})
t.Run("newtoncg success", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, presetBits := presetGrad(t, theta, presetVals)
// Σ (z − θ)² minimises at z = θ, where the gradient vanishes.
objective := func(z *Tensor) (*Tensor, error) {
d, err := z.Sub(theta)
if err != nil {
return nil, err
}
sq, err := d.Mul(d)
if err != nil {
return nil, err
}
return sq.Sum()
}
x0, err := core.FromFloats([]float64{0, 0}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
got, fv, err := MinimiseNewtonCG(objective, x0, NewtonCGOptions{})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
for i := range 2 {
if math.Abs(got.FloatAt(i)-thetaVals[i]) > 1e-8 {
t.Errorf("minimiser[%d] = %v, want %v", i, got.FloatAt(i), thetaVals[i])
}
}
if math.Abs(fv) > 1e-12 {
t.Errorf("value = %v, want 0", fv)
}
requireGradUntouched(t, "MinimiseNewtonCG", theta, preset, presetBits)
})
t.Run("newtoncg error", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, presetBits := presetGrad(t, theta, presetVals)
x0, err := core.FromFloats([]float64{1}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
if _, _, err := MinimiseNewtonCG(constant(theta), x0, NewtonCGOptions{}); err == nil {
t.Fatal("expected the disconnected-objective error")
}
requireGradUntouched(t, "MinimiseNewtonCG error path", theta, preset, presetBits)
})
}
// presetGrad installs vals as theta's accumulated gradient and
// returns the array the caller set together with a byte-exact snapshot
// of its payload, so the check below can prove both that the very array
// survived and that nothing wrote through it.
func presetGrad(t *testing.T, theta *Tensor, vals []float64) (*core.Array, []uint64) {
t.Helper()
g, err := core.FromFloats(vals, theta.Data().Len())
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
theta.SetGrad(g)
return g, gradBits(g)
}
// gradBits copies the raw float bits of a gradient array.
func gradBits(a *core.Array) []uint64 {
fs := a.RawFloats()
out := make([]uint64, len(fs))
for i, v := range fs {
out[i] = math.Float64bits(v)
}
return out
}
// requireGradUntouched requires the preset gradient to be the very
// array that was set and to hold the exact bits the snapshot captured:
// the helper under test must not write into it, replace it or clear it.
func requireGradUntouched(t *testing.T, label string, theta *Tensor, want *core.Array, before []uint64) {
t.Helper()
got := theta.Grad()
if got == nil {
t.Errorf("%s: the caller's gradient was cleared", label)
return
}
if got != want {
t.Errorf("%s: the caller's gradient array was replaced", label)
}
after := gradBits(got)
if len(after) != len(before) {
t.Errorf("%s: the caller's gradient changed length, %d to %d", label, len(before), len(after))
return
}
for i := range before {
if after[i] != before[i] {
t.Errorf("%s: gradient[%d] = %v, want the preset bits of %v", label, i,
got.FloatAt(i), math.Float64frombits(before[i]))
}
}
}
// TestBroadcastToRefusesIntOperands pins the dtype gate on BroadcastTo,
// the one shape op that used to accept an int tensor and record a graph
// node for it. The shape is validated first, so an impossible target
// keeps reporting the mismatch (the probe below asserts the same
// order).
func TestBroadcastToRefusesIntOperands(t *testing.T) {
src, err := core.FromInts([]int64{1, 2, 3}, 3)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
if out, err := FromArray(src, true).BroadcastTo(2, 3); err == nil {
t.Errorf("BroadcastTo accepted an int tensor: shape %v dtype %s",
out.Data().Shape(), out.Data().Dtype())
} else if !strings.Contains(err.Error(), "needs a float") {
t.Errorf("int refusal = %v, want the dtype error", err)
}
one, err := core.FromInts([]int64{7}, 1)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
if _, err := FromArray(one, true).BroadcastTo(2, 3); err == nil {
t.Error("BroadcastTo accepted an int (1,) tensor")
}
square, err := core.FromInts([]int64{1, 2, 3, 4}, 2, 2)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
if _, err := FromArray(square, true).BroadcastTo(2, 3); err == nil {
t.Error("an impossible broadcast target was accepted")
} else if !strings.Contains(err.Error(), "cannot broadcast") {
t.Errorf("impossible int broadcast = %v, want the shape refusal", err)
}
// A float tensor still broadcasts and differentiates; a complex one
// remains inside the accepted dtypes.
xf, err := core.FromFloats([]float64{1, 2, 3}, 3)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(xf, true)
y, err := xt.BroadcastTo(2, 3)
if err != nil {
t.Fatalf("float BroadcastTo: %v", err)
}
loss, err := y.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
for i := range xt.Grad().Len() {
if got := xt.Grad().FloatAt(i); got != 2 {
t.Errorf("float broadcast gradient[%d] = %v, want 2", i, got)
}
}
zc, err := core.FromComplexes([]complex128{1 + 1i, 2 - 0.5i}, 2)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
zt := FromArray(zc, true)
cz, err := zt.BroadcastTo(3, 2)
if err != nil {
t.Fatalf("complex BroadcastTo: %v", err)
}
csq, err := cz.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
closs, err := csq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := closs.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// L = Σ|broadcast(z)|² = 3·Σ|z|², so dL/dz̄ = 3z.
for i := range 2 {
if got, want := zt.Grad().ComplexAt(i), 3*zc.ComplexAt(i); got != want {
t.Errorf("complex broadcast gradient[%d] = %v, want %v", i, got, want)
}
}
}
// TestComplexMeanOfEmptyIsAnError pins the degenerate reduction: the
// complex branch of Mean divides by the element count, so an empty
// tensor used to answer 0/0 = NaN while the real branch errors loudly.
func TestComplexMeanOfEmptyIsAnError(t *testing.T) {
ec, err := core.FromComplexes([]complex128{}, 0)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
m, cerr := FromArray(ec, true).Mean()
if cerr == nil {
t.Fatalf("complex Mean of an empty tensor returned %v instead of an error", m.Data().ComplexAt(0))
}
if !strings.Contains(cerr.Error(), "empty array has no mean") {
t.Errorf("complex Mean error = %v, want the empty-reduction refusal", cerr)
}
er, err := core.FromFloats([]float64{}, 0)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
_, rerr := FromArray(er, true).Mean()
if rerr == nil {
t.Fatal("real Mean of an empty tensor was accepted")
}
// The two dtypes answer with one message.
if rerr.Error() != cerr.Error() {
t.Errorf("messages disagree: real %q, complex %q", rerr, cerr)
}
// A non-empty complex mean still reduces and differentiates.
z, err := core.FromComplexes([]complex128{1 + 1i, 3 - 1i}, 2)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
zt := FromArray(z, true)
mean, err := zt.Mean()
if err != nil {
t.Fatalf("Mean: %v", err)
}
if got, want := mean.Data().ComplexAt(0), complex(2, 0); got != want {
t.Errorf("complex mean = %v, want %v", got, want)
}
loss, err := mean.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// L = |mean|², so dL/dz̄ = mean/2 per element.
for i := range 2 {
if got, want := zt.Grad().ComplexAt(i), complex(1, 0); got != want {
t.Errorf("gradient[%d] = %v, want %v", i, got, want)
}
}
}
// TestAbs2KeepsFloat32Width pins Abs2's real branch: a float32 operand
// squares in float64 and stays float32, exactly as Pow(2) does,
// instead of promoting the forward result to float64.
func TestAbs2KeepsFloat32Width(t *testing.T) {
vals := []float32{1.3, -2.7, 0.5}
a, err := core.FromFloat32s(vals, 3)
if err != nil {
t.Fatalf("FromFloat32s: %v", err)
}
xt := FromArray(a, true)
sq, err := xt.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
if got := sq.Data().Dtype(); got != core.Float32 {
t.Errorf("float32 Abs2 output dtype = %s, want float32", got)
}
pw, err := xt.Pow(2)
if err != nil {
t.Fatalf("Pow: %v", err)
}
if got := pw.Data().Dtype(); got != core.Float32 {
t.Errorf("float32 Pow(2) output dtype = %s, want float32", got)
}
for i, v := range vals {
want := float32(float64(v) * float64(v))
if got := sq.Data().FloatAt(i); got != float64(want) {
t.Errorf("Abs2[%d] = %v, want the once-rounded %v", i, got, want)
}
if got := pw.Data().FloatAt(i); got != float64(want) {
t.Errorf("Pow(2)[%d] = %v, want %v", i, got, want)
}
}
// L = Σx² has dL/dx = 2x, and the float32 leaf keeps its width.
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if got := g.Dtype(); got != core.Float32 {
t.Errorf("float32 leaf gradient dtype = %s, want float32", got)
}
for i, v := range vals {
if got, want := g.FloatAt(i), 2*float64(v); math.Abs(got-want) > 1e-6 {
t.Errorf("gradient[%d] = %v, want %v", i, got, want)
}
}
// A float64 operand is untouched: float64 out, exact squares.
af, err := core.FromFloats([]float64{1.3, -2.7}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
sqf, err := FromArray(af, true).Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
if got := sqf.Data().Dtype(); got != core.Float {
t.Errorf("float64 Abs2 output dtype = %s, want float", got)
}
for i := range 2 {
if got, want := sqf.Data().FloatAt(i), af.FloatAt(i)*af.FloatAt(i); got != want {
t.Errorf("float64 Abs2[%d] = %v, want %v", i, got, want)
}
}
}
// TestNewtonCGConstantObjectiveHitsDisconnectedGuard covers
// MinimiseNewtonCG's g == nil guard, which the suite missed: its
// "constant objective" case slices the point itself, so a gradient
// exists and the guard never fires. A constant built from an
// independent graduated tensor (the disconnected-objective pattern) has no path to
// the starting point at all.
func TestNewtonCGConstantObjectiveHitsDisconnectedGuard(t *testing.T) {
c, err := FromFloat64s([]float64{3}, true, 1)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
constant := func(*Tensor) (*Tensor, error) { return c.Mul(c) }
x0, err := core.FromFloats([]float64{1}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
if _, _, err := MinimiseNewtonCG(constant, x0, NewtonCGOptions{MaxIterations: 2}); err == nil {
t.Fatal("a genuinely constant objective minimised without error")
} else if !strings.Contains(err.Error(), "does not depend on the starting point") {
t.Fatalf("error = %v, want the disconnected-graph refusal", err)
}
}
// TestSampleHMCLeavesCallerGradients pins SampleHMC's internal gradient
// evaluations: one leapfrog step runs one reverse pass, so a density
// closing over a graduated tensor used to add a contribution per step
// and one per proposal, on the success and the error path alike. The
// pass commits nothing, so the closed-over tensor keeps its accumulated
// gradient bit for bit, the guarantee Hessian, HessianVectorProduct,
// MinimiseNewtonCG and AdjointODE document.
func TestSampleHMCLeavesCallerGradients(t *testing.T) {
thetaVals := []float64{2, 1.5}
presetVals := []float64{0.5, -0.25}
thetaArr, err := core.FromFloats(thetaVals, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
q0, err := core.FromFloats([]float64{0.5, -0.5}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
// log π(q) = −½·Σ θ_i q_i²: a density that closes over a graduated
// tensor the caller owns, so any committed pass shows up in θ.
density := func(theta *Tensor) func(*Tensor) (*Tensor, error) {
return func(q *Tensor) (*Tensor, error) {
sq, err := q.Mul(q)
if err != nil {
return nil, err
}
w, err := sq.Mul(theta)
if err != nil {
return nil, err
}
s, err := w.Sum()
if err != nil {
return nil, err
}
return s.Scale(-0.5)
}
}
t.Run("success", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, bits := presetGrad(t, theta, presetVals)
calls := 0
counted := func(q *Tensor) (*Tensor, error) {
calls++
return density(theta)(q)
}
samples, err := SampleHMC(counted, q0,
HMCOptions{Step: 0.1, Steps: 2, Samples: 2, Thin: 1, Seed: 11})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
if got := samples.Shape(); got[0] != 2 || got[1] != 2 {
t.Fatalf("samples shape %v, want [2 2]", got)
}
// One evaluation at q0 plus one per leapfrog step: the run really
// did differentiate the density several times.
if calls < 3 {
t.Fatalf("SampleHMC ran %d gradient evaluations, want at least 3", calls)
}
requireGradUntouched(t, "SampleHMC", theta, preset, bits)
})
t.Run("midChainError", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, bits := presetGrad(t, theta, presetVals)
// The first evaluation at q0 succeeds, so a reverse pass has run;
// every proposal then reports the state as outside the support,
// which rejects the trajectory instead of aborting the run.
calls, refused := 0, 0
failing := func(q *Tensor) (*Tensor, error) {
calls++
if calls > 1 {
refused++
return nil, errf("outside the support")
}
return density(theta)(q)
}
samples, err := SampleHMC(failing, q0,
HMCOptions{Step: 0.1, Steps: 3, Samples: 1, Thin: 1, Seed: 12})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
if got := samples.Shape(); got[0] != 1 || got[1] != 2 {
t.Fatalf("samples shape %v, want [1 2]", got)
}
if refused == 0 {
t.Fatal("the density was never forced to fail mid-chain")
}
requireGradUntouched(t, "SampleHMC mid-chain error", theta, preset, bits)
})
t.Run("startError", func(t *testing.T) {
theta := FromArray(thetaArr, true)
preset, bits := presetGrad(t, theta, presetVals)
fails := func(*Tensor) (*Tensor, error) { return nil, errf("no density at the start") }
if _, err := SampleHMC(fails, q0,
HMCOptions{Step: 0.1, Steps: 2, Samples: 1, Thin: 1, Seed: 13}); err == nil {
t.Fatal("expected the start-time density error to be fatal")
}
requireGradUntouched(t, "SampleHMC start error", theta, preset, bits)
})
}
+230
View File
@@ -0,0 +1,230 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"strings"
"testing"
ode "sourcedock.dev/petrbalvin/tensor/integrate"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression tests for gradient hygiene: the Add backward
// handed both inputs the same gradient instance, AdjointODE polluted
// trainable leaves closed over by f but outside params, and
// MinimiseNewtonCG panicked on nil inputs where the rest of the
// package returns errors.
// TestAddBackwardIndependentGradients pins the aliasing fix: the two
// inputs of an Add receive two independent gradient buffers with
// identical values, so a write through one leaf's gradient cannot
// corrupt the other's.
func TestAddBackwardIndependentGradients(t *testing.T) {
x, err := FromFloat64s([]float64{2, 3}, true, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
y, err := FromFloat64s([]float64{4, 5}, true, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
z, err := x.Add(y)
if err != nil {
t.Fatalf("Add: %v", err)
}
loss, err := z.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
if x.Grad() == y.Grad() {
t.Fatal("Add handed both leaves the same gradient instance")
}
// d(x+y)/dx = 1 and d(x+y)/dy = 1, element for element.
for i := range 2 {
if x.Grad().FloatAt(i) != 1 {
t.Fatalf("x.grad[%d] = %v, want 1", i, x.Grad().FloatAt(i))
}
if y.Grad().FloatAt(i) != 1 {
t.Fatalf("y.grad[%d] = %v, want 1", i, y.Grad().FloatAt(i))
}
}
// A write through one leaf's buffer must leave the other's intact.
x.Grad().SetFloatAt(0, 999)
if y.Grad().FloatAt(0) != 1 {
t.Fatalf("y.grad[0] = %v after a write through x's gradient, want 1",
y.Grad().FloatAt(0))
}
}
// TestAddSameLeafAccumulatesTwice pins Add(x, x): the same leaf as both
// inputs accumulates both contributions into one gradient of 2.
func TestAddSameLeafAccumulatesTwice(t *testing.T) {
x, err := FromFloat64s([]float64{0.5, -1.25}, true, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
z, err := x.Add(x)
if err != nil {
t.Fatalf("Add: %v", err)
}
loss, err := z.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// d(x+x)/dx = 2, and the accumulated buffer is one array.
for i := range 2 {
if x.Grad().FloatAt(i) != 2 {
t.Fatalf("x.grad[%d] = %v, want 2", i, x.Grad().FloatAt(i))
}
}
}
// TestAdjointODELeavesHiddenLeavesClean pins the vjp fix: a trainable
// leaf closed over by f but not listed in params receives no gradient,
// because the Jacobian-vector products come from a pass that commits
// nothing.
func TestAdjointODELeavesHiddenLeavesClean(t *testing.T) {
theta, err := FromFloat64s([]float64{0.7}, true, 1)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
hidden, err := FromFloat64s([]float64{1.5}, true, 1)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(t float64, y *Tensor) (*Tensor, error) {
rate, err := theta.Neg()
if err != nil {
return nil, err
}
hm, err := y.Mul(hidden)
if err != nil {
return nil, err
}
return hm.Mul(rate)
}
y0, _ := core.FromFloats([]float64{1}, 1)
seed, _ := core.FromFloats([]float64{1}, 1)
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, y0, seed,
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
// dL/dθ = −1.5·e^{−1.05}, the central-difference answer the exact
// dynamics give.
want := -1.5 * math.Exp(-1.05)
if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 {
t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want)
}
if hidden.Grad() != nil {
t.Fatalf("the hidden leaf's gradient = %v, want nil", hidden.Grad())
}
if theta.Grad() != nil {
t.Fatalf("the parameter's gradient = %v, want nil", theta.Grad())
}
}
// TestAdjointODEThetaAgainstCentralDifferences checks the returned
// parameter sensitivity of the same hidden-leaf system against central
// differences of the forward solve, so the cleanup left the θ gradient
// every bit as accurate as it was.
func TestAdjointODEThetaAgainstCentralDifferences(t *testing.T) {
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
hidden, _ := FromFloat64s([]float64{1.5}, true, 1)
f := func(t float64, y *Tensor) (*Tensor, error) {
rate, err := theta.Neg()
if err != nil {
return nil, err
}
hm, err := y.Mul(hidden)
if err != nil {
return nil, err
}
return hm.Mul(rate)
}
y0, _ := core.FromFloats([]float64{1}, 1)
seed, _ := core.FromFloats([]float64{1}, 1)
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, y0, seed,
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
if err != nil {
t.Fatalf("AdjointODE: %v", err)
}
// Central differences on the loss y(1), with the parameter leaf's
// data swapped out for the perturbed values.
forwardLoss := func() float64 {
end, ferr := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) {
out, oerr := f(t, FromArray(ya, false))
if oerr != nil {
return nil, oerr
}
return out.Data(), nil
}, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15})
if ferr != nil {
t.Fatalf("forward solve: %v", ferr)
}
return end.FloatAt(0)
}
const eps = 1e-6
orig := theta.Data().FloatAt(0)
up, _ := core.FromFloats([]float64{orig + eps}, 1)
theta.ReplaceWith(up)
hi := forwardLoss()
dn, _ := core.FromFloats([]float64{orig - eps}, 1)
theta.ReplaceWith(dn)
lo := forwardLoss()
back, _ := core.FromFloats([]float64{orig}, 1)
theta.ReplaceWith(back)
fd := (hi - lo) / (2 * eps)
if math.Abs(blocks[0].FloatAt(0)-fd) > 1e-5*math.Max(1, math.Abs(fd)) {
t.Fatalf("dL/dθ: adjoint %.10g, central difference %.10g",
blocks[0].FloatAt(0), fd)
}
}
// TestMinimiseNewtonCGNilInputs pins the validation contract: a nil
// objective and a nil starting point are errors naming the argument,
// in step with SampleHMC, never panics.
func TestMinimiseNewtonCGNilInputs(t *testing.T) {
f := func(z *Tensor) (*Tensor, error) { return z.Sum() }
if _, _, err := MinimiseNewtonCG(nil, nil, NewtonCGOptions{}); err == nil {
t.Fatal("expected an error for a nil objective")
} else if !strings.Contains(err.Error(), "must not be nil") {
t.Fatalf("error = %v, want a must-not-be-nil refusal", err)
}
x0, err := core.FromFloats([]float64{1}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
if _, _, err := MinimiseNewtonCG(nil, x0, NewtonCGOptions{}); err == nil {
t.Fatal("expected an error for a nil objective")
} else if !strings.Contains(err.Error(), "f must not be nil") {
t.Fatalf("error = %v, want a refusal naming f", err)
}
if _, _, err := MinimiseNewtonCG(f, nil, NewtonCGOptions{}); err == nil {
t.Fatal("expected an error for a nil starting point")
} else if !strings.Contains(err.Error(), "starting point must not be nil") {
t.Fatalf("error = %v, want a refusal naming the starting point", err)
}
}
// TestODETraceSingleNodeRefused pins the interpolation guard: a trace
// with fewer than two recorded nodes has no interval to interpolate
// over, so the accessor refuses instead of indexing out of range.
func TestODETraceSingleNodeRefused(t *testing.T) {
tr := &odeTrace{times: []float64{1}, states: [][]float64{{2, 3}},
slopes: [][]float64{{0, 0}}, dim: 2}
if _, err := tr.at(1); err == nil {
t.Fatal("expected an error for a one-node trace")
} else if !strings.Contains(err.Error(), "at least two") {
t.Fatalf("error = %v, want a refusal naming the node count", err)
}
}
+47
View File
@@ -0,0 +1,47 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// numericGrad estimates the gradient of a scalar function at a by
// central differences, the reference the analytic backward is checked
// against.
func numericGrad(f func(a *core.Array) float64, a *core.Array) []float64 {
n := a.Len()
out := make([]float64, n)
for i := range n {
hi := 1e-6
up := cloneFlat(a)
down := cloneFlat(a)
up[i] += hi
down[i] -= hi
au, _ := core.FromFloats(up, a.Shape()...)
ad, _ := core.FromFloats(down, a.Shape()...)
out[i] = (f(au) - f(ad)) / (2 * hi)
}
return out
}
func cloneFlat(a *core.Array) []float64 {
out := make([]float64, a.Len())
for i := range a.Len() {
out[i] = a.FloatAt(i)
}
return out
}
func maxAbsDiff(a, b []float64) float64 {
m := 0.0
for i := range a {
if d := math.Abs(a[i] - b[i]); d > m {
m = d
}
}
return m
}
+178
View File
@@ -0,0 +1,178 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Second-order differentiation, forward-over-reverse: the
// inner derivative is the exact analytic gradient Backward produces,
// and only the outer derivative runs by central differences over one
// coordinate at a time. The result carries the accuracy of the exact
// first derivative with the O(h²) truncation of the outer stencil, the
// same trade a hand-written finite-difference Hessian makes but with
// none of the first-order error.
// HessianOptions tunes Hessian and HessianVectorProduct. Step is the
// absolute coordinate perturbation (≤ 0 picks sqrt(eps)·max(1, |x_i|)
// per coordinate, the stencil that balances truncation against
// cancellation at double precision).
type HessianOptions struct {
Step float64
}
// Hessian returns the Hessian matrix of a scalar function f at x, an
// (n, n) float64 array for an n-element x. f receives a tensor that
// requires grad and must return a single-element real tensor; complex
// outputs are rejected like Backward does. The cost is 2n gradient
// evaluations, the price of a dense second derivative by any method
// that does not exploit structure; for large n prefer
// HessianVectorProduct. The evaluations differentiate the graph
// without committing anything, so the accumulated gradients of the
// tensors f closes over are left exactly as they were, on the success
// and the error path alike.
func Hessian(f func(*Tensor) (*Tensor, error), x *Tensor, opts HessianOptions) (*core.Array, error) {
if x.Data().Dtype() == core.Complex {
return nil, base.Errf("Hessian: complex points are not supported")
}
n := x.Data().Len()
if n == 0 {
return nil, base.Errf("Hessian: the point must not be empty")
}
p := flatFloats(x.Data())
out := zeros(core.Float, []int{n, n})
h := opts.Step
for j := range n {
step := h
if step <= 0 {
step = math.Sqrt(2.220446049250313e-16) * math.Max(1, math.Abs(p[j]))
}
gp, err := hessianColumn(f, x, p, j, step)
if err != nil {
return nil, err
}
gm, err := hessianColumn(f, x, p, j, -step)
if err != nil {
return nil, err
}
inv2h := 1 / (2 * step)
for i := range n {
out.SetFloatAt(i*n+j, (gp[i]-gm[i])*inv2h)
}
}
return out, nil
}
// hessianColumn evaluates the analytic gradient of f at the point
// perturbed by step along coordinate j, flattened.
func hessianColumn(f func(*Tensor) (*Tensor, error), x *Tensor, p []float64, j int, step float64) ([]float64, error) {
probe := append([]float64(nil), p...)
probe[j] += step
pa, err := core.FromFloats(probe, x.Data().Shape()...)
if err != nil {
return nil, base.Errf("Hessian: %w", err)
}
xt := FromArray(pa, true)
y, err := f(xt)
if err != nil {
return nil, base.Errf("Hessian: %w", err)
}
if y.Data().Len() != 1 {
return nil, base.Errf("Hessian: f must return a scalar, got %d elements", y.Data().Len())
}
grads, err := y.reverseGrads()
if err != nil {
return nil, base.Errf("Hessian: %w", err)
}
g := grads[xt]
if g == nil {
return nil, base.Errf("Hessian: the objective does not depend on x, so no gradient exists")
}
return flatFloats(g), nil
}
// HessianVectorProduct returns H·v, the Hessian of the scalar f at x
// contracted with the direction v, by a central difference along the
// direction itself, with the step scaled so it never depends on v's
// magnitude. Two gradient evaluations
// answer for any n, which is what makes Newton-CG tractable where a
// dense Hessian is not. As in Hessian, the evaluations leave the
// accumulated gradients of every tensor f closes over untouched.
func HessianVectorProduct(f func(*Tensor) (*Tensor, error), x, v *Tensor, opts HessianOptions) (*core.Array, error) {
if x.Data().Dtype() == core.Complex {
return nil, base.Errf("HessianVectorProduct: complex points are not supported")
}
if v.Data().Dtype() == core.Complex {
return nil, base.Errf("HessianVectorProduct: complex directions are not supported")
}
n := x.Data().Len()
if v.Data().Len() != n {
return nil, base.Errf("HessianVectorProduct: direction has %d elements for %d variables",
v.Data().Len(), n)
}
vn := 0.0
for _, v := range flatFloats(v.Data()) {
vn += v * v
}
vn = math.Sqrt(vn)
if vn == 0 {
// H·0 = 0 in the shape of the point, the same shape the
// quotient below returns: a flat vector here would change the
// result's shape with the direction's norm.
return zeros(core.Float, x.Data().Shape()), nil
}
h := opts.Step
if h <= 0 {
h = 1e-5
}
p := flatFloats(x.Data())
vf := flatFloats(v.Data())
eval := func(sign float64) ([]float64, error) {
probe := make([]float64, n)
for i := range n {
probe[i] = p[i] + sign*h*vf[i]/vn
}
pa, err := core.FromFloats(probe, x.Data().Shape()...)
if err != nil {
return nil, base.Errf("HessianVectorProduct: %w", err)
}
xt := FromArray(pa, true)
y, err := f(xt)
if err != nil {
return nil, base.Errf("HessianVectorProduct: %w", err)
}
if y.Data().Len() != 1 {
return nil, base.Errf("HessianVectorProduct: f must return a scalar, got %d elements", y.Data().Len())
}
grads, err := y.reverseGrads()
if err != nil {
return nil, base.Errf("HessianVectorProduct: %w", err)
}
g := grads[xt]
if g == nil {
return nil, base.Errf("HessianVectorProduct: the objective does not depend on x, so no gradient exists")
}
return flatFloats(g), nil
}
gp, err := eval(1)
if err != nil {
return nil, err
}
gm, err := eval(-1)
if err != nil {
return nil, err
}
// The step advanced h·v/|v| along v, so the quotient is (H·v/|v|)
// and carries the |v| factor back in.
out := zeros(core.Float, x.Data().Shape())
inv2h := vn / (2 * h)
for i := range n {
out.SetFloatAt(i, (gp[i]-gm[i])*inv2h)
}
return out, nil
}
+40
View File
@@ -0,0 +1,40 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"strings"
"testing"
)
// Regression tests: an objective whose graph
// never reaches x left xt.Grad() nil, and the second-order helpers
// dereferenced it.
// TestHessianDisconnectedObjective pins the error: an objective that
// ignores its argument has no gradient to differentiate.
func TestHessianDisconnectedObjective(t *testing.T) {
x, err := FromFloat64s([]float64{1, 2}, true, 2)
if err != nil {
t.Fatal(err)
}
c, err := FromFloat64s([]float64{3}, true, 1)
if err != nil {
t.Fatal(err)
}
objective := func(*Tensor) (*Tensor, error) { return c.Mul(c) }
if _, err := Hessian(objective, x, HessianOptions{}); err == nil {
t.Fatal("expected an error when the objective does not depend on x")
} else if !strings.Contains(err.Error(), "does not depend") {
t.Fatalf("error = %v, want a disconnected-graph refusal", err)
}
v, err := FromFloat64s([]float64{1, 0}, true, 2)
if err != nil {
t.Fatal(err)
}
if _, err := HessianVectorProduct(objective, x, v, HessianOptions{}); err == nil {
t.Fatal("expected an error when the objective does not depend on x")
}
}
+193
View File
@@ -0,0 +1,193 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// TestHessianQuadratic pins the exact case: for f(x) = ½xᵀAx + bᵀx the
// Hessian is A whatever the point.
func TestHessianQuadratic(t *testing.T) {
a := []float64{4, 1, 1, 3}
b := []float64{-1, 2}
x0 := []float64{0.5, -1.25}
xt, err := FromFloat64s(x0, false, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
az, err := FromFloat64s(a, false, 2, 2)
if err != nil {
return nil, err
}
bz, err := FromFloat64s(b, false, 2)
if err != nil {
return nil, err
}
halfAz, err := az.Scale(0.5)
if err != nil {
return nil, err
}
azx, err := halfAz.MatMul(z)
if err != nil {
return nil, err
}
sum1, err := azx.Add(bz)
if err != nil {
return nil, err
}
// ½xᵀAx + bᵀx = ((½A)x + b)·x
return sum1.Mul(z)
}
// ((½A)x + b)·x is elementwise; the scalar loss needs the sum.
fScalar := func(z *Tensor) (*Tensor, error) {
p, err := f(z)
if err != nil {
return nil, err
}
return p.Sum()
}
h, err := Hessian(fScalar, xt, HessianOptions{})
if err != nil {
t.Fatalf("Hessian: %v", err)
}
for i := range 2 {
for j := range 2 {
if math.Abs(h.FloatAt(i*2+j)-a[i*2+j]) > 1e-6 {
t.Fatalf("H[%d][%d] = %g, want %g", i, j, h.FloatAt(i*2+j), a[i*2+j])
}
}
}
}
// TestHessianRosenbrock pins a nonquadratic landscape against the
// analytic Hessian of the 2-D Rosenbrock function.
func TestHessianRosenbrock(t *testing.T) {
x0 := []float64{-0.5, 1.25}
xt, err := FromFloat64s(x0, false, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
els := []int{0, 1}
x0t, err := z.Slice(0, els[0], els[0]+1)
if err != nil {
return nil, err
}
x1t, err := z.Slice(0, els[1], els[1]+1)
if err != nil {
return nil, err
}
x0sq, err := x0t.Pow(2)
if err != nil {
return nil, err
}
diff, err := x1t.Sub(x0sq)
if err != nil {
return nil, err
}
term1v, err := diff.Pow(2)
if err != nil {
return nil, err
}
one, err := FromFloat64s([]float64{1}, false, 1)
if err != nil {
return nil, err
}
x0m1, err := x0t.Sub(one)
if err != nil {
return nil, err
}
term2v, err := x0m1.Pow(2)
if err != nil {
return nil, err
}
term2s, err := term2v.Scale(100)
if err != nil {
return nil, err
}
total, err := term1v.Add(term2s)
if err != nil {
return nil, err
}
return total.Sum()
}
h, err := Hessian(f, xt, HessianOptions{})
if err != nil {
t.Fatalf("Hessian: %v", err)
}
x, y := x0[0], x0[1]
// Analytic Hessian of f = (y − x²)² + 100(x − 1)².
h00 := 12*x*x - 4*y + 200
h01 := -4 * x
h11 := 2.0
want := [][]float64{{h00, h01}, {h01, h11}}
for i := range 2 {
for j := range 2 {
if math.Abs(h.FloatAt(i*2+j)-want[i][j]) > 1e-4*math.Max(1, math.Abs(want[i][j])) {
t.Fatalf("H[%d][%d] = %g, want %g", i, j, h.FloatAt(i*2+j), want[i][j])
}
}
}
}
// TestHessianVectorProduct pins H·v against the dense Hessian.
func TestHessianVectorProduct(t *testing.T) {
a := []float64{4, 1, 1, 3}
x0 := []float64{0.5, -1.25}
xt, err := FromFloat64s(x0, false, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
az, err := FromFloat64s(a, false, 2, 2)
if err != nil {
return nil, err
}
halfAz, err := az.Scale(0.5)
if err != nil {
return nil, err
}
azx, err := halfAz.MatMul(z)
if err != nil {
return nil, err
}
p, err := azx.Mul(z)
if err != nil {
return nil, err
}
return p.Sum()
}
vArr, err := core.FromFloats([]float64{2, -1}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
v := FromArray(vArr, false)
hv, err := HessianVectorProduct(f, xt, v, HessianOptions{})
if err != nil {
t.Fatalf("HessianVectorProduct: %v", err)
}
// A·v exactly.
want := []float64{4*2 + 1*(-1), 1*2 + 3*(-1)}
for i := range 2 {
if math.Abs(hv.FloatAt(i)-want[i]) > 1e-5 {
t.Fatalf("Hv[%d] = %g, want %g", i, hv.FloatAt(i), want[i])
}
}
}
// TestHessianRejectsVectorOutput pins the scalar contract.
func TestHessianRejectsVectorOutput(t *testing.T) {
xt, err := FromFloat64s([]float64{1, 2}, false, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(z *Tensor) (*Tensor, error) { return z, nil }
if _, err := Hessian(f, xt, HessianOptions{}); err == nil {
t.Fatal("Hessian accepted a vector output")
}
}
+194
View File
@@ -0,0 +1,194 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Hamiltonian Monte Carlo on the autograd surface. The target is any
// differentiable log density: a trajectory simulates the Hamiltonian
// dynamics of a unit-mass particle on that landscape, with the
// momentum refreshed from a standard Gaussian each round and the
// leapfrog integrator driven by gradients Backward computes, so the
// user supplies the density and the chain does the calculus. The
// Metropolis correction on the trajectory's energy change makes the
// stationary distribution exact despite the integration error.
// HMCOptions tunes SampleHMC. Step is the leapfrog step size and
// Steps the number of leapfrog steps per trajectory, both
// problem-dependent with no sensible default; BurnIn trajectories are
// discarded before every Thin-th trajectory contributes one sample,
// until Samples have been collected, and a Thin of zero or less
// quietly normalises to one. Seed feeds the package's own
// xoshiro generator, so a run is bit-reproducible.
type HMCOptions struct {
Step float64
Steps int
BurnIn int
Thin int
Samples int
Seed int64
}
// SampleHMC draws Samples states from the unnormalised density whose
// logarithm logDensity computes, starting at the vector q0 and
// returning the kept states as a (Samples × dim) array; rejected
// trajectories repeat the current state, as Markov chain sampling
// does. logDensity receives a leaf tensor requiring grad and must
// return a scalar tensor connected to it; an error it raises at q0 is
// fatal, while one raised inside a proposal marks the state as
// outside the support and rejects the trajectory, which is how
// constrained densities keep the chain away from forbidden regions.
// A nil density, a non-vector start, a non-positive step, step count
// or sample count, or a density that does not yield a gradient are
// errors. The gradient evaluations differentiate the graph without
// committing anything, so the accumulated gradients of the tensors
// logDensity closes over are left exactly as they were, on the success
// and the error path alike.
func SampleHMC(logDensity func(q *Tensor) (*Tensor, error),
q0 *core.Array, opts HMCOptions) (*core.Array, error) {
const name = "SampleHMC"
if logDensity == nil {
return nil, errf("%s: logDensity must not be nil", name)
}
if q0 == nil {
return nil, errf("%s: the state must not be nil", name)
}
if q0.NDim() != 1 || q0.Len() == 0 {
return nil, errf("%s: the state must be a non-empty vector, got shape %s",
name, prettyShape(q0.Shape()))
}
if q0.Dtype() == core.Complex {
return nil, errf("%s: complex states are not supported", name)
}
if opts.Step <= 0 {
return nil, errf("%s: Step must be positive, got %g", name, opts.Step)
}
if opts.Steps <= 0 {
return nil, errf("%s: Steps must be at least 1, got %d", name, opts.Steps)
}
if opts.Samples <= 0 {
return nil, errf("%s: Samples must be at least 1, got %d", name, opts.Samples)
}
if opts.BurnIn < 0 {
return nil, errf("%s: BurnIn must not be negative, got %d", name, opts.BurnIn)
}
if opts.Thin <= 0 {
opts.Thin = 1
}
dim := q0.Len()
rng := core.NewGenerator(opts.Seed)
// eval computes log π at q together with ∇log π(q): a fresh leaf
// per call, one reverse pass that commits nothing, plain floats out.
// A leapfrog trajectory calls this once per step, so a pass that
// committed would add one contribution to every tensor logDensity
// closes over per step; the reverse pass of reverseGrads leaves the
// caller's accumulated gradients untouched instead.
eval := func(q []float64) (float64, []float64, error) {
data, err := core.FromFloats(q, dim)
if err != nil {
return 0, nil, errf("%s: %w", name, err)
}
leaf := FromArray(data, true)
out, err := logDensity(leaf)
if err != nil {
return 0, nil, err
}
if out.Data().Len() != 1 {
return 0, nil, errf("%s: logDensity returned shape %s, want a scalar",
name, prettyShape(out.Data().Shape()))
}
grads, err := out.reverseGrads()
if err != nil {
return 0, nil, errf("%s: %w", name, err)
}
g := grads[leaf]
if g == nil || g.Len() != dim {
return 0, nil, errf("%s: logDensity did not yield a gradient of length %d", name, dim)
}
return out.Data().FloatAt(0), flatFloats(g), nil
}
q := flatFloats(q0)
logPi, gradient, err := eval(q)
if err != nil {
return nil, errf("%s: %w", name, err)
}
p := make([]float64, dim)
qNew := make([]float64, dim)
pNew := make([]float64, dim)
values := make([]float64, 0, opts.Samples*dim)
trajectories := opts.BurnIn + opts.Samples*opts.Thin
for traj := 1; traj <= trajectories; traj++ {
// Fresh momentum from the standard Gaussian; the kinetic
// energy is p·p/2 for unit mass.
momentum, err := core.Normal(rng, dim, 0, 1)
if err != nil {
return nil, errf("%s: %w", name, err)
}
kinetic0 := 0.0
momentumF := flatFloats(momentum)
for i := range p {
p[i] = momentumF[i]
kinetic0 += p[i] * p[i] / 2
}
copy(qNew, q)
copy(pNew, p)
// Leapfrog: half kick, drift, full kick per step, one gradient
// evaluation each, the last one doubling as the proposal's
// log density.
diverged := false
logPiNew := math.Inf(-1)
g := gradient
for range opts.Steps {
for i := range pNew {
pNew[i] += opts.Step / 2 * g[i]
}
for i := range qNew {
qNew[i] += opts.Step * pNew[i]
}
lp, gNew, eerr := eval(qNew)
if eerr != nil {
diverged = true
break
}
for i := range pNew {
pNew[i] += opts.Step / 2 * gNew[i]
}
g = gNew
logPiNew = lp
}
if !diverged {
kinetic1 := 0.0
for i := range pNew {
kinetic1 += pNew[i] * pNew[i] / 2
}
// Metropolis on the energy change; an undefined change
// (a NaN crept into the landscape) rejects.
logAccept := (-logPi + kinetic0) - (-logPiNew + kinetic1)
uniform, uerr := core.Floats(rng, 1)
if uerr != nil {
return nil, errf("%s: %w", name, uerr)
}
if !math.IsNaN(logAccept) && math.Log(uniform.FloatAt(0)) < logAccept {
copy(q, qNew)
logPi = logPiNew
gradient = g
}
}
if traj <= opts.BurnIn || (traj-opts.BurnIn)%opts.Thin != 0 {
continue
}
values = append(values, q...)
}
samples, err := core.FromFloats(values, opts.Samples, dim)
if err != nil {
return nil, errf("%s: %w", name, err)
}
return samples, nil
}
+217
View File
@@ -0,0 +1,217 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// gaussianLogDensity builds the log density of independent standard
// normals: log π(q) = −‖q‖²/2.
func gaussianLogDensity(q *Tensor) (*Tensor, error) {
sq, err := q.Mul(q)
if err != nil {
return nil, err
}
total, err := sq.Sum()
if err != nil {
return nil, err
}
return total.Scale(-0.5)
}
// gammaLogDensity builds log π(x) = log x − x, the unnormalised log
// density of a Gamma(2, 1) distribution.
func gammaLogDensity(q *Tensor) (*Tensor, error) {
logq, err := q.Log()
if err != nil {
return nil, err
}
shifted, err := logq.Sub(q)
if err != nil {
return nil, err
}
return shifted.Sum()
}
// TestSampleHMCNormal runs the chain on a two-dimensional standard
// normal and pins the sample moments: the target's mean is zero, its
// variance one and its components independent. A fixed seed makes the
// draw deterministic, so the bounds are checked facts about this run,
// not hopes about a random one.
func TestSampleHMCNormal(t *testing.T) {
q0, err := core.FromFloats([]float64{2, -2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
samples, err := SampleHMC(gaussianLogDensity, q0,
HMCOptions{Step: 0.3, Steps: 20, BurnIn: 500, Samples: 6000, Thin: 1, Seed: 42})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
if got := samples.Shape(); got[0] != 6000 || got[1] != 2 {
t.Fatalf("samples shape %v, want [6000 2]", got)
}
rows := samples.Shape()[0]
mean := []float64{0, 0}
variance := []float64{0, 0}
for r := range rows {
for c := range 2 {
v := samples.FloatAt(r*2 + c)
mean[c] += v / float64(rows)
}
}
for r := range rows {
for c := range 2 {
d := samples.FloatAt(r*2+c) - mean[c]
variance[c] += d * d / float64(rows)
}
}
for c := range 2 {
if math.Abs(mean[c]) > 0.15 {
t.Fatalf("mean[%d] = %.4g, want |mean| ≤ 0.15", c, mean[c])
}
if math.Abs(variance[c]-1) > 0.2 {
t.Fatalf("variance[%d] = %.4g, want 1 ± 0.2", c, variance[c])
}
}
// Cross moment of the independent components.
cov := 0.0
for r := range rows {
cov += (samples.FloatAt(r*2) - mean[0]) * (samples.FloatAt(r*2+1) - mean[1]) / float64(rows)
}
if math.Abs(cov) > 0.15 {
t.Fatalf("cross moment = %.4g, want |cov| ≤ 0.15", cov)
}
}
// TestSampleHMCDeterministic checks the seed contract: the same seed
// replays bit-identically, a different seed does not.
func TestSampleHMCDeterministic(t *testing.T) {
q0, _ := core.FromFloats([]float64{1}, 1)
run := func(seed int64) *core.Array {
samples, err := SampleHMC(gaussianLogDensity, q0,
HMCOptions{Step: 0.4, Steps: 16, BurnIn: 100, Samples: 200, Seed: seed})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
return samples
}
a, b := run(7), run(7)
c := run(8)
for i := range a.Len() {
if a.FloatAt(i) != b.FloatAt(i) {
t.Fatalf("the same seed produced different samples at %d", i)
}
}
same := true
for i := range a.Len() {
if a.FloatAt(i) != c.FloatAt(i) {
same = false
break
}
}
if same {
t.Fatal("different seeds produced identical samples")
}
}
// TestSampleHMCSupportedDensity samples a Gamma(2, 1) target,
// log π(x) = log x − x on x > 0. Proposals that overshoot into the
// forbidden half-line yield a NaN density and are rejected, so every
// kept sample stays positive and the mean approaches 2.
func TestSampleHMCSupportedDensity(t *testing.T) {
q0, _ := core.FromFloats([]float64{1}, 1)
samples, err := SampleHMC(gammaLogDensity, q0,
HMCOptions{Step: 0.3, Steps: 10, BurnIn: 500, Samples: 6000, Seed: 3})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
sum := 0.0
for i := range samples.Len() {
x := samples.FloatAt(i)
if x <= 0 {
t.Fatalf("sample %d = %g left the support", i, x)
}
sum += x
}
mean := sum / float64(samples.Len())
if math.Abs(mean-2) > 0.15 {
t.Fatalf("Gamma(2) mean = %.4g, want 2 ± 0.15", mean)
}
}
// TestSampleHMCRejectedByGuard exercises the explicit rejection path:
// the density errors outside its support instead of returning NaN,
// and the chain still stays inside it.
func TestSampleHMCRejectedByGuard(t *testing.T) {
q0, _ := core.FromFloats([]float64{0.5}, 1)
samples, err := SampleHMC(func(q *Tensor) (*Tensor, error) {
x := q.Data().FloatAt(0)
if x <= 0 {
return nil, errf("outside the support")
}
return gammaLogDensity(q)
}, q0, HMCOptions{Step: 0.5, Steps: 20, BurnIn: 300, Samples: 2000, Seed: 5})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
for i := range samples.Len() {
if samples.FloatAt(i) <= 0 {
t.Fatalf("sample %d = %g left the support", i, samples.FloatAt(i))
}
}
}
// TestSampleHMCErrors pins the validation contract, including a
// density that fails at the start, returns a non-scalar, or never
// touches the leaf core.
func TestSampleHMCErrors(t *testing.T) {
q0, _ := core.FromFloats([]float64{1}, 1)
if _, err := SampleHMC(nil, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a nil density")
}
rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2)
if _, err := SampleHMC(gaussianLogDensity, rank2, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
if _, err := SampleHMC(gaussianLogDensity, nil, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a nil state")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a non-positive step")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Samples: 1}); err == nil {
t.Fatal("expected an error for a non-positive step count")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5}); err == nil {
t.Fatal("expected an error for a non-positive sample count")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1, BurnIn: -3}); err == nil {
t.Fatal("expected an error for a negative BurnIn")
}
nonScalar := func(q *Tensor) (*Tensor, error) {
two, _ := core.FromFloats([]float64{1, 2}, 2)
return FromArray(two, false), nil
}
if _, err := SampleHMC(nonScalar, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a non-scalar density")
}
detached := func(q *Tensor) (*Tensor, error) {
one, _ := core.FromFloats([]float64{1}, 1)
return FromArray(one, false), nil
}
if _, err := SampleHMC(detached, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a density disconnected from the leaf")
}
failsAtStart := func(q *Tensor) (*Tensor, error) {
return nil, errf("no density at the start")
}
if _, err := SampleHMC(failsAtStart, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected the start-time density error to be fatal")
}
}
+38
View File
@@ -0,0 +1,38 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import "testing"
// TestHessianVectorProductZeroDirectionShape pins the repair: a zero
// direction returned a flat length-n vector while
// every other answer carries the point's own shape.
func TestHessianVectorProductZeroDirectionShape(t *testing.T) {
x0 := []float64{0.5, -1.25, 2, -2}
xt, err := FromFloat64s(x0, false, 2, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(z *Tensor) (*Tensor, error) { return z.Sum() }
v, err := FromFloat64s(make([]float64, 4), false, 2, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
hv, err := HessianVectorProduct(f, xt, v, HessianOptions{})
if err != nil {
t.Fatalf("HessianVectorProduct: %v", err)
}
want := []int{2, 2}
got := hv.Shape()
for i := range want {
if got[i] != want[i] {
t.Fatalf("H·0 shape = %v, want %v", got, want)
}
}
for i := range 4 {
if hv.FloatAt(i) != 0 {
t.Fatalf("H·0 = %v, want the zero matrix", hv.FloatAt(i))
}
}
}
+158
View File
@@ -0,0 +1,158 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression pins for the MatMul adjoints: the 1-D × 2-D branch against
// central differences, real and complex, and the Newton-CG loop's last
// iteration.
// TestMatMulVectorByMatrixGradient pins the 1-D × 2-D branch of
// the MatMul adjoint (da = g·Bᵀ, db = outer(a, g)), real and complex,
// against central differences. The branch had no gradient test at all.
func TestMatMulVectorByMatrixGradient(t *testing.T) {
avec := []float64{1.5, -0.5, 2, 0.25}
bvec := []float64{
0.5, -1, 2,
1.5, 0.25, -0.75,
-2, 1, 0.5,
1, -0.5, 1.25,
}
// A weighted linear loss, so every output slot carries its own
// coefficient and a wrong routing cannot cancel against another.
w := []float64{0.7, -1.3, 2.1}
lossOf := func(a, b *core.Array) float64 {
out, err := core.MatMul2D(a, b)
if err != nil {
t.Fatalf("MatMul2D: %v", err)
}
s := 0.0
for i := range out.Len() {
s += w[i] * out.FloatAt(i)
}
return s
}
cloneWith := func(a *core.Array, i int, v float64) *core.Array {
vals := make([]float64, a.Len())
for k := range a.Len() {
vals[k] = a.FloatAt(k)
}
vals[i] = v
out, err := core.FromFloats(vals, a.Shape()...)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
return out
}
a, _ := FromFloat64s(avec, true, 4)
b, _ := FromFloat64s(bvec, true, 4, 3)
if err := backwardWeightedMatMul(t, a, b, w); err != nil {
t.Fatalf("Backward: %v", err)
}
baseA, _ := core.FromFloats(avec, 4)
baseB, _ := core.FromFloats(bvec, 4, 3)
for i := range 4 {
eps := 1e-6
want := (lossOf(cloneWith(baseA, i, avec[i]+eps), baseB) -
lossOf(cloneWith(baseA, i, avec[i]-eps), baseB)) / (2 * eps)
if math.Abs(a.Grad().FloatAt(i)-want) > 1e-5*(1+math.Abs(want)) {
t.Fatalf("da[%d] = %g, want %g", i, a.Grad().FloatAt(i), want)
}
}
for i := range 12 {
eps := 1e-6
want := (lossOf(baseA, cloneWith(baseB, i, bvec[i]+eps)) -
lossOf(baseA, cloneWith(baseB, i, bvec[i]-eps))) / (2 * eps)
if math.Abs(b.Grad().FloatAt(i)-want) > 1e-5*(1+math.Abs(want)) {
t.Fatalf("db[%d] = %g, want %g", i, b.Grad().FloatAt(i), want)
}
}
// The complex 1-D × 2-D branch against numericComplexGrad, under
// the same weighted fold backwardComplex builds: L = Σ Re(w̄·y) +
// Σ|y|²/n with the helper's own deterministic weights.
cb := FromArray(mustComplexes([]complex128{
0.5 + 0.5i, -1,
1.5, 0.25 - 0.75i,
-2 + 1i, 1,
}, 3, 2), true)
op := func(x *Tensor) (*Tensor, error) { return x.MatMul(cb) }
vals := []complex128{1 + 0.5i, -0.25 - 1i, 0.75 + 0.25i}
xt := backwardComplex(t, op, vals, 3)
closs := complexLossOf(t, op, spectralWeights(2, 11))
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(closs, xt.Data()), 1e-6)
}
// backwardWeightedMatMul builds L = w·(a·B) over fresh leaves and runs
// one backward pass.
func backwardWeightedMatMul(t *testing.T, a, b *Tensor, w []float64) error {
t.Helper()
out, err := a.MatMul(b)
if err != nil {
return err
}
wv, err := FromFloat64s(w, false, len(w))
if err != nil {
return err
}
loss, err := out.Mul(wv)
if err != nil {
return err
}
sum, err := loss.Sum()
if err != nil {
return err
}
sum.Backward()
return nil
}
// TestNewtonCGConvergesOnTheLastIteration pins that a tolerance
// met exactly on the final permitted iteration is a success, not a
// budget error whose message prints a gradient already under the
// tolerance.
func TestNewtonCGConvergesOnTheLastIteration(t *testing.T) {
f := func(x *Tensor) (*Tensor, error) {
d, err := x.Sub(mustTensorF64(1.5))
if err != nil {
return nil, err
}
return d.Mul(d)
}
x0, _ := core.FromFloats([]float64{0}, 1)
// One truncated-CG step lands within the Hessian-product rounding
// of the minimiser (about 1e-11 here); a tolerance of 1e-9 is met
// by exactly that step, on the final permitted iteration.
out, _, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{MaxIterations: 1, Tolerance: 1e-9})
if err != nil {
t.Fatalf("MinimiseNewtonCG on a quadratic with one exact step: %v", err)
}
if math.Abs(out.FloatAt(0)-1.5) > 1e-9 {
t.Fatalf("minimum = %g, want 1.5", out.FloatAt(0))
}
}
// mustTensorF64 wraps one float as a no-grad tensor.
func mustTensorF64(v float64) *Tensor {
t, err := FromFloat64s([]float64{v}, false, 1)
if err != nil {
panic(err)
}
return t
}
// mustComplexes builds a complex array or fails the test.
func mustComplexes(vals []complex128, shape ...int) *core.Array {
a, err := core.FromComplexes(vals, shape...)
if err != nil {
panic(err)
}
return a
}
+248
View File
@@ -0,0 +1,248 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Newton-CG minimisation: truncated conjugate gradients on
// the Hessian system, driven by autograd. It lives in the grad package
// because it is meaningless without the graph: the gradients come from
// Backward and the Hessian never forms, each CG iteration buying one
// Hessian-vector product for two backward passes. That is the
// optimiser large problems want, where a dense second derivative does
// not fit memory and the numerical-difference optimisers of the optim
// package lose their accuracy.
// NewtonCGOptions tunes MinimiseNewtonCG. MaxIterations bounds the
// outer Newton steps (default 100); Tolerance stops when the gradient
// norm falls under it (default 1e-8); MaxCGIterations bounds the inner
// CG solve per outer step (default n, the problem dimension).
type NewtonCGOptions struct {
MaxIterations int
Tolerance float64
MaxCGIterations int
}
// MinimiseNewtonCG returns the point and value of a local minimum of
// the scalar objective f near x0 by the Newton-CG method: each step
// solves H·s = −∇f with truncated conjugate gradients (negative
// curvature stops the solve and falls back to the first direction),
// then an Armijo backtracking line search secures descent. f receives
// a leaf tensor and must return a single-element real tensor. A
// non-finite objective, an unreachable Armijo condition or an
// exhausted iteration budget is an error naming the state it stopped
// in; the converged answer is a fresh array the caller owns. The
// gradient evaluations differentiate the graph without committing
// anything, so the accumulated gradients of the tensors f closes over
// are left exactly as they were, on the success and the error path
// alike.
func MinimiseNewtonCG(f func(*Tensor) (*Tensor, error), x0 *core.Array, opts NewtonCGOptions) (*core.Array, float64, error) {
const name = "MinimiseNewtonCG"
if f == nil {
return nil, 0, errf("%s: f must not be nil", name)
}
if x0 == nil {
return nil, 0, errf("%s: the starting point must not be nil", name)
}
n := x0.Len()
if n == 0 {
return nil, 0, errf("%s: the starting point must have at least one element", name)
}
if x0.Dtype() == core.Complex {
return nil, 0, errf("%s: complex starting points are not supported", name)
}
maxIter := opts.MaxIterations
if maxIter <= 0 {
maxIter = 100
}
tol := opts.Tolerance
if tol <= 0 {
tol = 1e-8
}
maxCG := opts.MaxCGIterations
if maxCG <= 0 {
maxCG = n
}
// eval runs the objective and its reverse pass at point p, returning
// the loss and the flattened gradient. The pass commits nothing, so
// the caller's own gradients survive the evaluation untouched.
eval := func(p *core.Array) (float64, *core.Array, error) {
xt := FromArray(p, true)
y, err := f(xt)
if err != nil {
return 0, nil, errf("%s: %w", name, err)
}
if y.Data().Len() != 1 {
return 0, nil, errf("%s: the objective must return a scalar, got %d elements", name, y.Data().Len())
}
v := y.Data().FloatAt(0)
if math.IsNaN(v) || math.IsInf(v, 0) {
return 0, nil, errf("%s: the objective is non-finite (%g)", name, v)
}
grads, err := y.reverseGrads()
if err != nil {
return 0, nil, errf("%s: %w", name, err)
}
g := grads[xt]
if g == nil {
return 0, nil, errf("%s: the objective does not depend on the starting point", name)
}
return v, g, nil
}
x := x0
f0, g, err := eval(x)
if err != nil {
return nil, 0, err
}
for iter := 1; iter <= maxIter; iter++ {
gnorm := flatNorm(g)
if gnorm <= tol {
return clonePoint(x), f0, nil
}
// Truncated CG on H·s = −g. The Hessian acts through the
// Hessian-vector product, two backward passes per iteration.
xt := FromArray(x, true)
s := make([]float64, n)
r := make([]float64, n)
p := make([]float64, n)
gs := make([]float64, n)
gFloats := flatFloats(g)
for i := range n {
r[i] = -gFloats[i]
p[i] = r[i]
gs[i] = gFloats[i]
}
rr := 0.0
for i := range n {
rr += r[i] * r[i]
}
for cg := 0; cg < maxCG; cg++ {
pArr, herr := core.FromFloats(p, n)
if herr != nil {
return nil, 0, errf("%s: %w", name, herr)
}
hp, herr2 := HessianVectorProduct(f, xt, FromArray(pArr, false), HessianOptions{})
if herr2 != nil {
return nil, 0, errf("%s: %w", name, herr2)
}
hpF := flatFloats(hp)
pHp := 0.0
for i := range n {
pHp += p[i] * hpF[i]
}
if pHp <= 0 {
// Negative or vanishing curvature: the quadratic model
// is not convex here. The first iteration falls back to
// the steepest descent direction; later ones keep what
// the solve has accumulated.
if cg == 0 {
copy(s, p)
}
break
}
alpha := rr / pHp
for i := range n {
s[i] += alpha * p[i]
r[i] -= alpha * hpF[i]
}
rrNew := 0.0
for i := range n {
rrNew += r[i] * r[i]
}
if math.Sqrt(rrNew) <= 0.1*gnorm {
break
}
beta := rrNew / rr
for i := range n {
p[i] = r[i] + beta*p[i]
}
rr = rrNew
}
// Armijo backtracking along s; gᵀs is negative by construction.
gsDot := 0.0
for i := range n {
gsDot += gs[i] * s[i]
}
if gsDot >= 0 {
return nil, 0, errf("%s: the CG direction does not descend at step %d", name, iter)
}
step := 1.0
xFloats := flatFloats(x)
var xNew *core.Array
var fNew float64
accepted := false
for range 40 {
vals := make([]float64, n)
for i := range n {
vals[i] = xFloats[i] + step*s[i]
}
cand, cerr := core.FromFloats(vals, x.Shape()...)
if cerr != nil {
return nil, 0, errf("%s: %w", name, cerr)
}
// cand assigns to the outer xNew; a := here would shadow
// it and hand the post-loop update a nil.
fNew, g, err = eval(cand)
if err != nil {
return nil, 0, err
}
if fNew <= f0+1e-4*step*gsDot {
accepted = true
xNew = cand
break
}
step /= 2
}
if !accepted {
return nil, 0, errf("%s: the line search found no descent at step %d (f = %.6g)", name, iter, f0)
}
x = xNew
f0 = fNew
}
// The last accepted step updated g after the loop-top test, so a
// run whose tolerance was met exactly on the final iteration must
// re-test before the budget refusal reports it; the message below
// would otherwise print a gradient already under the tolerance.
if flatNorm(g) <= tol {
return clonePoint(x), f0, nil
}
return nil, 0, errf("%s: no convergence in %d steps (gradient norm %.3g)", name, maxIter, flatNorm(g))
}
// flatNorm returns the Euclidean norm of a flattened gradient. The sum
// runs in ascending element order on the raw payload when it can; the
// walk is bounded by the element count, not the payload, because a
// rebased view's storage may run longer than its own elements.
func flatNorm(g *core.Array) float64 {
s := 0.0
if !g.Strided() && g.Dtype() == core.Float {
gs := g.RawFloats()
for i := range g.Len() {
s += gs[i] * gs[i]
}
return math.Sqrt(s)
}
for i := range g.Len() {
s += g.FloatAt(i) * g.FloatAt(i)
}
return math.Sqrt(s)
}
// clonePoint copies the converged point so the caller owns it.
func clonePoint(x *core.Array) *core.Array {
vals := make([]float64, x.Len())
for i := range vals {
vals[i] = x.FloatAt(i)
}
out, _ := core.FromFloats(vals, x.Shape()...)
return out
}
+197
View File
@@ -0,0 +1,197 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// TestNewtonCGQuadratic pins the exactly-Newtonian case: a quadratic
// with SPD Hessian converges to the analytic minimiser in a couple of
// steps.
func TestNewtonCGQuadratic(t *testing.T) {
a := []float64{4, 1, 1, 3}
b := []float64{-1, 2}
x0, err := core.FromFloats([]float64{0.5, -1.25}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
az, err := FromFloat64s(a, false, 2, 2)
if err != nil {
return nil, err
}
bz, err := FromFloat64s(b, false, 2)
if err != nil {
return nil, err
}
halfA, err := az.Scale(0.5)
if err != nil {
return nil, err
}
azx, err := halfA.MatMul(z)
if err != nil {
return nil, err
}
lin, err := azx.Add(bz)
if err != nil {
return nil, err
}
prod, err := lin.Mul(z)
if err != nil {
return nil, err
}
return prod.Sum()
}
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-12})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
// x* = −A⁻¹b: solve 4x+y = 1, x+3y = −2 so x = 5/11, y = −9/11.
if math.Abs(x.FloatAt(0)-5.0/11) > 1e-9 || math.Abs(x.FloatAt(1)+9.0/11) > 1e-9 {
t.Fatalf("minimiser = (%g, %g), want (5/11, -9/11)", x.FloatAt(0), x.FloatAt(1))
}
// f* = ½x*ᵀAx* + bᵀx* = 253/242 − 23/11 = −253/242.
const want = -253.0 / 242.0
if math.Abs(fv-want) > 1e-10 {
t.Fatalf("value = %.12g, want %.12g", fv, want)
}
}
// TestNewtonCGRosenbrock pins a nonquadratic valley: the classic
// Rosenbrock minimum at (1, 1) from the far side.
func TestNewtonCGRosenbrock(t *testing.T) {
x0, err := core.FromFloats([]float64{-1.5, 2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
x0t, err := z.Slice(0, 0, 1)
if err != nil {
return nil, err
}
x1t, err := z.Slice(0, 1, 2)
if err != nil {
return nil, err
}
x0sq, err := x0t.Pow(2)
if err != nil {
return nil, err
}
diff, err := x1t.Sub(x0sq)
if err != nil {
return nil, err
}
term1, err := diff.Pow(2)
if err != nil {
return nil, err
}
one, err := FromFloat64s([]float64{1}, false, 1)
if err != nil {
return nil, err
}
x0m1, err := x0t.Sub(one)
if err != nil {
return nil, err
}
term2, err := x0m1.Pow(2)
if err != nil {
return nil, err
}
term2s, err := term2.Scale(100)
if err != nil {
return nil, err
}
total, err := term1.Add(term2s)
if err != nil {
return nil, err
}
return total.Sum()
}
x, _, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-7, MaxIterations: 200})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
if math.Abs(x.FloatAt(0)-1) > 1e-4 || math.Abs(x.FloatAt(1)-1) > 1e-4 {
t.Fatalf("minimiser = (%.6f, %.6f), want (1, 1)", x.FloatAt(0), x.FloatAt(1))
}
}
// TestNewtonCGNegativeCurvature pins the fallback: a double well
// whose start sits in the concave region between the minima. The CG
// must take its steepest-descent fallback there and still land in a
// well.
func TestNewtonCGNegativeCurvature(t *testing.T) {
x0, err := core.FromFloats([]float64{0.1, 0.2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
// f = Σ(x⁴ − x²): Hessian 12x² − 2 is negative for |x| < 1/√6,
// so the start is concave; the wells sit at ±1/√2 per coordinate.
f := func(z *Tensor) (*Tensor, error) {
q, err := z.Pow(4)
if err != nil {
return nil, err
}
sq, err := z.Abs2()
if err != nil {
return nil, err
}
d, err := q.Sub(sq)
if err != nil {
return nil, err
}
return d.Sum()
}
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-9})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
const well = 1.0 / math.Sqrt2
for i := range 2 {
if math.Abs(math.Abs(x.FloatAt(i))-well) > 1e-6 {
t.Fatalf("coordinate %d = %g, want magnitude %g", i, x.FloatAt(i), well)
}
}
// f at a well: Σ(1/4 − 1/2) = −1/2.
if math.Abs(fv+0.5) > 1e-9 {
t.Fatalf("value = %.12g, want -0.5", fv)
}
}
// TestNewtonCGScalarInput pins the n = 1 path and the error contract.
func TestNewtonCGScalarInput(t *testing.T) {
x0, err := core.FromFloats([]float64{3}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
sq, err := z.Pow(2)
if err != nil {
return nil, err
}
four, err := sq.Scale(4)
if err != nil {
return nil, err
}
return four.Sum()
}
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
if math.Abs(x.FloatAt(0)) > 1e-7 || math.Abs(fv) > 1e-12 {
t.Fatalf("minimiser = %g, value = %g", x.FloatAt(0), fv)
}
// A slice of the point itself: the objective is disconnected from
// the minimised point, so the run exhausts its iterations on a
// constant value and errors loudly.
c, _ := core.FromFloats([]float64{1}, 1)
if _, _, err := MinimiseNewtonCG(func(z *Tensor) (*Tensor, error) { return z.Slice(0, 0, 1) }, c, NewtonCGOptions{}); err == nil {
t.Fatal("a slice of the point itself minimised without error")
}
}
+42
View File
@@ -0,0 +1,42 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// TransposeAxes reorders the axes of a tensor by dims. The backward
// applies the inverse permutation to the incoming gradient: axis moves
// are invertible data motion, so no element mixing occurs and the
// gradient is exactly the same move played backwards.
func (t *Tensor) TransposeAxes(dims ...int) (*Tensor, error) {
if err := t.checkDiff("TransposeAxes"); err != nil {
return nil, err
}
out, err := core.TransposeAxes(t.data, dims...)
if err != nil {
return nil, err
}
perm := append([]int(nil), dims...)
orig := t.data.Shape()
return t.unaryResult("TransposeAxes", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
dx, err := core.TransposeAxes(g.arr, inversePerm(perm)...)
if err != nil {
return err
}
dst[0] = gradSlot{arr: dx, sh: orig}
return nil
}), nil
}
// inversePerm flips an axis permutation: if out = permute(x, p), then
// permute(out, p⁻¹) restores x's axis order.
func inversePerm(perm []int) []int {
inv := make([]int, len(perm))
for i, p := range perm {
inv[p] = i
}
return inv
}
+100
View File
@@ -0,0 +1,100 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
func TestTensorTransposeAxesValues(t *testing.T) {
x, _ := core.FromFloats([]float64{
1, 2, 3,
4, 5, 6,
}, 2, 3)
out, err := FromArray(x, false).TransposeAxes(1, 0)
if err != nil {
t.Fatal(err)
}
if got := out.Data().Shape(); got[0] != 3 || got[1] != 2 {
t.Fatalf("shape: %v", got)
}
want := []float64{1, 4, 2, 5, 3, 6}
for i := range want {
if g := out.Data().FloatAt(i); g != want[i] {
t.Fatalf("[%d] = %v, want %v", i, g, want[i])
}
}
// A rank-3 rotation moves the trailing axis to the front.
y, _ := core.FromFloats([]float64{
1, 2, 3, 4,
5, 6, 7, 8,
9, 10, 11, 12,
}, 2, 3, 2)
rotated, err := FromArray(y, false).TransposeAxes(2, 0, 1)
if err != nil {
t.Fatal(err)
}
if got := rotated.Data().Shape(); got[0] != 2 || got[1] != 2 || got[2] != 3 {
t.Fatalf("rank-3 shape: %v", got)
}
// Invalid permutations error before any graph work.
if _, err := FromArray(x, false).TransposeAxes(0, 0); err == nil {
t.Fatal("duplicate axis accepted")
}
if _, err := FromArray(x, false).TransposeAxes(0); err == nil {
t.Fatal("short permutation accepted")
}
}
// TestTensorTransposeAxesGradient routes a weighted sum through the
// permutation: the analytic input gradient is exactly the weight tensor
// played back through the inverse permutation.
func TestTensorTransposeAxesGradient(t *testing.T) {
x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
w, _ := core.FromFloats([]float64{0.5, -1, 2, 0.25, -0.75, 1.5}, 3, 2)
xt := FromArray(x, true)
joint, err := xt.TransposeAxes(1, 0) // gives (3, 2)
if err != nil {
t.Fatal(err)
}
scaled, err := joint.Mul(FromArray(w, false))
if err != nil {
t.Fatal(err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
g := xt.Grad()
if g == nil || g.Dtype() != core.Float {
t.Fatalf("gradient missing or wrong dtype: %v", g)
}
for i := range 6 {
row, col := i/3, i%3
if got := g.FloatAt(i); got != w.FloatAt(col*2+row) {
t.Errorf("grad[%d] = %v, want %v", i, got, w.FloatAt(col*2+row))
}
}
// Round-trip: permuting by (1,0) then back restores the values.
back, err := joint.TransposeAxes(1, 0)
if err != nil {
t.Fatal(err)
}
for i := range 6 {
if back.Data().FloatAt(i) != x.FloatAt(i) {
t.Fatalf("round-trip[%d] = %v, want %v", i, back.Data().FloatAt(i), x.FloatAt(i))
}
}
}
+352
View File
@@ -0,0 +1,352 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sync"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Gradient buffer recycling for the backward sweep. A sweep allocates
// one gradient array per node output and per folded contribution; the
// arrays die within the sweep that made them, except the ones committed
// to leaves, which escape to the caller. The pool reclaims the
// intermediates: a borrowed array arrives with a fully zeroed payload,
// the pool retains a bounded number of elements, and nothing is
// recycled while a live reference to it exists. That last rule rests on
// an invariant the closures maintain: a backward closure never returns
// the incoming gradient buffer itself, never returns one buffer in two
// slots and never returns a view of another live array, so an entry the
// sweep releases is unreachable from the graph, from the returned map
// and from every other entry.
//
// An array the pool declines (an exotic dtype, a rank above six, a
// payload above the cap, a full bucket) falls back to ordinary
// allocation; correctness never depends on a hit.
// gradSlot is one gradient array together with the shape it was built
// for, which is the pool's reuse key. Slots travel instead of bare
// arrays so the sweep can release a buffer without re-deriving its
// shape, which would cost an allocation of its own.
type gradSlot struct {
arr *core.Array
sh []int
}
// poolKey identifies a gradient buffer exactly: dtype, rank, element
// count and dimensions. The fixed dimension array keeps the key
// comparable, so buckets need no stored shape for matching.
type poolKey struct {
dt core.Dtype
nd int8
n int
d [6]int
}
// poolKeyOf builds the key for dt and shape, reporting false for
// anything the pool does not accept.
func poolKeyOf(dt core.Dtype, shape []int) (poolKey, bool) {
var k poolKey
if len(shape) == 0 || len(shape) > len(k.d) {
return k, false
}
switch dt {
case core.Float, core.Float32, core.Complex:
default:
return k, false
}
n := 1
for i, d := range shape {
n *= d
k.d[i] = d
}
if n > gradPoolMaxArrayElems {
return k, false
}
k.dt, k.nd, k.n = dt, int8(len(shape)), n
return k, true
}
// The retention caps: no more than gradPoolMaxArrayElems elements in
// one array, gradPoolMaxElems retained across all buckets and
// gradPoolPerBucket arrays of one exact shape. A buffer outside the
// caps is dropped to the garbage collector instead of retained, so the
// pool cannot pin memory beyond these bounds however hard one workload
// pushes it.
const (
gradPoolMaxArrayElems = 1 << 20
gradPoolMaxElems = 1 << 20
gradPoolPerBucket = 32
)
var gradPool = struct {
sync.Mutex
buckets map[poolKey][]*core.Array
elems int
}{buckets: make(map[poolKey][]*core.Array)}
// freshGrad allocates a zeroed array of k's dtype and shape, taking
// ownership of the payload the way the constructor documents: grad
// writes that payload only through the array's own raw accessor.
func freshGrad(k poolKey, shape []int) *core.Array {
var a *core.Array
switch k.dt {
case core.Float:
a, _ = core.FloatsFromArray(make([]float64, k.n), shape...)
case core.Float32:
a, _ = core.FromFloat32Slice(make([]float32, k.n), shape...)
default:
a, _ = core.ComplexFromArray(make([]complex128, k.n), shape...)
}
if a == nil {
a, _ = core.Zeros(k.dt, shape...)
}
return a
}
// clearGradPayload zeroes every slot of a's payload, the borrow-side
// rule: a recycled buffer must never carry the previous sweep's values
// into a reader.
func clearGradPayload(a *core.Array) {
switch a.Dtype() {
case core.Float:
clear(a.RawFloats())
case core.Float32:
clear(a.RawFloat32s())
case core.Complex:
clear(a.RawComplexes())
}
}
// releaseGrad offers a dead gradient array to the pool. The caller must
// have proven the array unreachable: releasing one a tape, a graph or a
// caller still holds would let the next borrower corrupt it. The shape
// must be the shape the array was built with; a mismatch is caught by
// the element-count check and drops the array instead of pooling it.
func releaseGrad(a *core.Array, shape []int) {
if a == nil || a.Strided() {
return
}
k, ok := poolKeyOf(a.Dtype(), shape)
if !ok || k.n != a.Len() {
return
}
gradPool.Lock()
b := gradPool.buckets[k]
if len(b) >= gradPoolPerBucket || gradPool.elems+k.n > gradPoolMaxElems {
gradPool.Unlock()
return
}
gradPool.buckets[k] = append(b, a)
gradPool.elems += k.n
gradPool.Unlock()
}
// gradArena is one sweep's private free list. A sweep borrows and
// releases in near-LIFO order, so most round trips stay on the calling
// goroutine under no lock; the global pool absorbs overflow and
// supplies misses, and a sweep-end flush returns the leftovers under a
// single lock. Every sweep owns its arena, so concurrent sweeps on
// different graphs never share one.
type gradArena struct {
free []gradFree
elems int
pooled bool
}
type gradFree struct {
arr *core.Array
sh []int
k poolKey
}
// gradArenaMaxFree bounds what one arena carries between flushes; a
// sweep that releases beyond it hands the surplus to the global pool.
const gradArenaMaxFree = 256
var gradArenaPool = sync.Pool{New: func() any { return &gradArena{} }}
func borrowArena() *gradArena {
ar := gradArenaPool.Get().(*gradArena)
ar.pooled = true
// A recycled arena comes back with its previous free list: the
// slice is emptied here, or stale entries would both starve the
// scan and pin the buffers they still name.
ar.free = ar.free[:0]
ar.elems = 0
return ar
}
// borrowGrad returns a zeroed array of dt and shape: the arena's own
// free list first, then the global pool, then fresh allocation. A nil
// arena means the legacy sweep path, which allocates exactly what it
// allocated before and takes no part in the pool.
func (ar *gradArena) borrowGrad(dt core.Dtype, shape []int) *core.Array {
if ar == nil {
a, _ := core.Zeros(dt, shape...)
return a
}
k, ok := poolKeyOf(dt, shape)
if !ok {
a, _ := core.Zeros(dt, shape...)
return a
}
for i := len(ar.free) - 1; i >= 0; i-- {
e := ar.free[i]
if e.k != k {
continue
}
ar.free[i] = ar.free[len(ar.free)-1]
ar.free = ar.free[:len(ar.free)-1]
ar.elems -= k.n
clearGradPayload(e.arr)
return e.arr
}
gradPool.Lock()
b := gradPool.buckets[k]
if len(b) > 0 {
a := b[len(b)-1]
gradPool.buckets[k] = b[:len(b)-1]
gradPool.elems -= k.n
gradPool.Unlock()
clearGradPayload(a)
return a
}
gradPool.Unlock()
return freshGrad(k, shape)
}
// releaseGrad returns a dead gradient array to the arena, falling
// through to the global pool when the arena is full. A nil arena means
// the caller owns the lifetime, so the array is left to the collector.
func (ar *gradArena) releaseGrad(a *core.Array, shape []int) {
if ar == nil || a == nil || a.Strided() {
return
}
k, ok := poolKeyOf(a.Dtype(), shape)
if !ok || k.n != a.Len() {
return
}
if len(ar.free) >= gradArenaMaxFree || ar.elems+k.n > gradPoolMaxElems {
releaseGrad(a, shape)
return
}
ar.free = append(ar.free, gradFree{arr: a, sh: shape, k: k})
ar.elems += k.n
}
// flush returns everything the arena still holds to the global pool.
// Buffers the pool declines are dropped to the collector; the arena
// itself returns to the sync.Pool for the next sweep.
func (ar *gradArena) flush() {
if ar == nil {
return
}
gradPool.Lock()
for _, e := range ar.free {
b := gradPool.buckets[e.k]
if len(b) >= gradPoolPerBucket || gradPool.elems+e.k.n > gradPoolMaxElems {
continue
}
gradPool.buckets[e.k] = append(b, e.arr)
gradPool.elems += e.k.n
}
gradPool.Unlock()
clear(ar.free)
ar.free = ar.free[:0]
ar.elems = 0
if ar.pooled {
gradArenaPool.Put(ar)
}
}
// fillGradSlotC writes z into every complex element of s.
func fillGradSlotC(s gradSlot, z complex128) {
if s.arr == nil {
return
}
cs := s.arr.RawComplexes()[:s.arr.Len()]
for i := range cs {
cs[i] = z
}
}
// fillGradSlot writes v into every element of s, the seed and fill
// helper. Each dtype takes the same spelling fillConst writes: a
// float32 destination narrows the constant once and stores it, a
// complex one stores complex(v, 0).
func fillGradSlot(s gradSlot, v float64) {
if s.arr == nil {
return
}
n := s.arr.Len()
switch s.arr.Dtype() {
case core.Float:
fs := s.arr.RawFloats()[:n]
for i := range fs {
fs[i] = v
}
case core.Float32:
fs := s.arr.RawFloat32s()[:n]
fv := float32(v)
for i := range fs {
fs[i] = fv
}
case core.Complex:
cs := s.arr.RawComplexes()[:n]
z := complex(v, 0)
for i := range cs {
cs[i] = z
}
}
}
// tapeFrame is one node's position in the reverse sweep's explicit
// walk: the node being expanded and the next operand index to visit.
type tapeFrame struct {
node *gradNode
next int
}
// tapeWork is the sweep's traversal scratch: the topological order, the
// walk stack and the seen set. It is borrowed per sweep and returned
// with its references cleared, so a pooled copy never pins a dead
// graph; a workload whose graph exceeds the retention cap drops the
// buffers to the collector instead of pinning them.
type tapeWork struct {
order []*gradNode
stack []tapeFrame
seen map[*Tensor]bool
}
const gradPoolMaxTapeNodes = 1 << 16
var tapeWorkPool = sync.Pool{New: func() any {
return &tapeWork{seen: make(map[*Tensor]bool)}
}}
func borrowTapeWork() *tapeWork {
w := tapeWorkPool.Get().(*tapeWork)
w.order = w.order[:0]
w.stack = w.stack[:0]
clear(w.seen)
if cap(w.order) > gradPoolMaxTapeNodes || cap(w.stack) > gradPoolMaxTapeNodes {
return &tapeWork{seen: make(map[*Tensor]bool)}
}
return w
}
func releaseTapeWork(w *tapeWork) {
if w == nil {
return
}
clear(w.order[:cap(w.order)])
clear(w.stack[:cap(w.stack)])
clear(w.seen)
if cap(w.order) > gradPoolMaxTapeNodes || cap(w.stack) > gradPoolMaxTapeNodes {
return
}
tapeWorkPool.Put(w)
}
+200
View File
@@ -0,0 +1,200 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"runtime"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The gradient pool's contract, measured: a repeated backward sweep on
// one process must keep the heap flat rather than growing with the
// iteration count, the sweep's arithmetic must be bit-for-bit
// reproducible across runs that share the pool, and the per-sweep cost
// itself is pinned by benchmarks that separate graph construction from
// the reverse pass.
// tapeChain builds a chain of n element-wise nodes over x and w and
// reduces it to a scalar, the fixture the sweep benchmarks repeat.
func tapeChain(t testing.TB, x, w *Tensor, n int) *Tensor {
t.Helper()
h := x
for i := range n {
var err error
switch i % 4 {
case 0:
h, err = h.Add(w)
case 1:
h, err = h.Mul(w)
case 2:
h, err = h.Tanh()
default:
h, err = h.Scale(0.25)
}
if err != nil {
t.Fatal(err)
}
}
s, err := h.Sum()
if err != nil {
t.Fatal(err)
}
return s
}
func tapeLeaf(t testing.TB, seed, n int) *Tensor {
t.Helper()
v := make([]float64, n)
for i := range v {
v[i] = 0.25 + float64((i*7+seed)%13)*0.125
}
a, err := core.FromFloats(v, n)
if err != nil {
t.Fatal(err)
}
return FromArray(a, true)
}
// TestGradPoolFlatHeap runs one repeated backward workload and checks
// the live heap stops growing: the pool's retention caps must bound
// what one process holds, however many sweeps it serves.
func TestGradPoolFlatHeap(t *testing.T) {
x, w := tapeLeaf(t, 1, 8), tapeLeaf(t, 2, 8)
s := tapeChain(t, x, w, 128)
run := func() {
for range 500 {
x.ZeroGrad()
w.ZeroGrad()
if err := s.Backward(); err != nil {
t.Fatal(err)
}
}
}
var early, late runtime.MemStats
runtime.GC()
run()
runtime.GC()
runtime.ReadMemStats(&early)
run()
run()
runtime.GC()
runtime.ReadMemStats(&late)
// Two more batches of a thousand sweeps may add pool slack but not
// a growth trend: the second reading stays within a small factor of
// the first, which a leaking pool would break.
if late.HeapInuse > early.HeapInuse*2+1<<20 {
t.Fatalf("heap grew across repeated sweeps: %d then %d bytes in use", early.HeapInuse, late.HeapInuse)
}
}
// TestGradBackwardDeterminismBits runs the same program twice through
// the pooled sweep and demands identical gradient bits: recycling a
// buffer must never leak a previous sweep's values into a result.
func TestGradBackwardDeterminismBits(t *testing.T) {
gradOf := func() []float64 {
x, w := tapeLeaf(t, 3, 8), tapeLeaf(t, 4, 8)
s := tapeChain(t, x, w, 64)
if err := s.Backward(); err != nil {
t.Fatal(err)
}
gx, gw := x.Grad(), w.Grad()
if gx == nil || gw == nil {
t.Fatal("missing leaf gradient")
}
out := make([]float64, 0, gx.Len()+gw.Len())
out = append(out, gx.RawFloats()[:gx.Len()]...)
out = append(out, gw.RawFloats()[:gw.Len()]...)
return out
}
// Warm the pool with unrelated sweeps, so the measured runs borrow
// recycled buffers carrying other work's values.
for range 64 {
a, b := tapeLeaf(t, 9, 8), tapeLeaf(t, 10, 8)
s := tapeChain(t, a, b, 32)
if err := s.Backward(); err != nil {
t.Fatal(err)
}
}
first, second := gradOf(), gradOf()
if len(first) != len(second) {
t.Fatalf("gradient lengths differ: %d and %d", len(first), len(second))
}
for i := range first {
if math.Float64bits(first[i]) != math.Float64bits(second[i]) {
t.Fatalf("gradient bit %d differs: %x and %x", i,
math.Float64bits(first[i]), math.Float64bits(second[i]))
}
}
}
// BenchmarkTapeChainSweepBackward measures the reverse sweep alone
// on a 129-node chain that is built once: every allocation here is the
// sweep's own, not the graph's.
func BenchmarkTapeChainSweepBackward(b *testing.B) {
x, w := tapeLeaf(b, 1, 8), tapeLeaf(b, 2, 8)
s := tapeChain(b, x, w, 128)
b.ReportAllocs()
for b.Loop() {
x.ZeroGrad()
w.ZeroGrad()
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkWaveTapeChainForwardRebuild measures building the same chain
// afresh with no backward pass, the per-node graph-construction cost
// the sweep benchmarks otherwise carry inside their loop.
func BenchmarkWaveTapeChainForwardRebuild(b *testing.B) {
x, w := tapeLeaf(b, 1, 8), tapeLeaf(b, 2, 8)
b.ReportAllocs()
for b.Loop() {
s := tapeChain(b, x, w, 128)
if s.Data().Len() != 1 {
b.Fatal("unexpected shape")
}
}
}
// BenchmarkWaveWideFanSweepBackward measures the reverse sweep of a
// 64-way fan over one shared leaf: the fold-heavy edge pattern, built
// once.
func BenchmarkWaveWideFanSweepBackward(b *testing.B) {
x := tapeLeaf(b, 3, 16)
leaves := make([]*Tensor, 64)
for i := range leaves {
leaves[i] = tapeLeaf(b, 10+i, 16)
}
acc, err := x.Mul(leaves[0])
if err != nil {
b.Fatal(err)
}
for _, l := range leaves[1:] {
p, err := x.Mul(l)
if err != nil {
b.Fatal(err)
}
if acc, err = acc.Add(p); err != nil {
b.Fatal(err)
}
}
s, err := acc.Sum()
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
x.ZeroGrad()
for _, l := range leaves {
l.ZeroGrad()
}
if err := s.Backward(); err != nil {
b.Fatal(err)
}
}
}
+164
View File
@@ -0,0 +1,164 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sync"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Pins for the pooled reverse sweep: the pooled Backward must answer
// bit-identically to the legacy map-returning sweep on the same graph,
// concurrent sweeps on separate graphs must answer the serial reference
// bits, and the pool's retention caps must refuse releases past them.
// pinLit builds a leaf of n elements from fixed literals, the fixture
// shape the tape benchmarks use.
func pinLit(seed, n int) *Tensor {
v := make([]float64, n)
for i := range v {
v[i] = 0.25 + float64((i*7+seed)%13)*0.125
}
a, err := core.FromFloats(v, n)
if err != nil {
panic(err)
}
return FromArray(a, true)
}
// chainLeafBits builds the same deep chain the tape benchmarks build,
// runs one pooled Backward and returns the two leaf gradients' raw
// bits. The chain fans both leaves into every node, so every fold
// accumulates multiple contributions.
func chainLeafBits(seedA, seedB, nodes int) ([]float64, error) {
x, w := pinLit(seedA, 8), pinLit(seedB, 8)
s, err := deepChain(x, w, nodes)
if err != nil {
return nil, err
}
x.ZeroGrad()
w.ZeroGrad()
if err := s.Backward(); err != nil {
return nil, err
}
xg, wg := x.Grad(), w.Grad()
if xg == nil || wg == nil {
return nil, errf("pinned chain: missing leaf gradient")
}
out := make([]float64, 0, 16)
out = append(out, xg.RawFloats()[:8]...)
out = append(out, wg.RawFloats()[:8]...)
return out, nil
}
func TestPooledBackwardBitsMatchLegacySweep(t *testing.T) {
pooled, err := chainLeafBits(1, 2, 64)
if err != nil {
t.Fatalf("pooled sweep: %v", err)
}
// The same graph through the legacy sweep: reverseGrads commits
// nothing, so its map carries the leaves' gradients from this pass
// alone, which is what the pooled sweep commits on a fresh leaf.
x, w := pinLit(1, 8), pinLit(2, 8)
s, err := deepChain(x, w, 64)
if err != nil {
t.Fatalf("legacy chain: %v", err)
}
grads, err := s.reverseGrads()
if err != nil {
t.Fatalf("legacy sweep: %v", err)
}
gx, gw := grads[x], grads[w]
if gx == nil || gw == nil {
t.Fatal("legacy sweep returned no leaf gradient")
}
legacy := append(append([]float64{}, gx.RawFloats()[:8]...), gw.RawFloats()[:8]...)
if len(legacy) != len(pooled) {
t.Fatalf("length %d, want %d", len(pooled), len(legacy))
}
for i := range pooled {
if math.Float64bits(pooled[i]) != math.Float64bits(legacy[i]) {
t.Fatalf("leaf gradient %d: pooled %#x, legacy %#x",
i, math.Float64bits(pooled[i]), math.Float64bits(legacy[i]))
}
}
}
func TestConcurrentBackwardDeterminism(t *testing.T) {
ref, err := chainLeafBits(3, 4, 48)
if err != nil {
t.Fatalf("serial reference: %v", err)
}
const sweeps = 40
outs := make([][]float64, 2)
errs := make([]error, 2)
var wg sync.WaitGroup
for g := range 2 {
wg.Go(func() {
for range sweeps {
bits, err := chainLeafBits(3, 4, 48)
if err != nil {
errs[g] = err
return
}
outs[g] = bits
}
})
}
wg.Wait()
for g := range 2 {
if errs[g] != nil {
t.Fatalf("goroutine %d: %v", g, errs[g])
}
for i := range ref {
if math.Float64bits(outs[g][i]) != math.Float64bits(ref[i]) {
t.Fatalf("goroutine %d element %d: %#x, want %#x",
g, i, math.Float64bits(outs[g][i]), math.Float64bits(ref[i]))
}
}
}
}
func TestGradPoolCapsAreEnforced(t *testing.T) {
k, ok := poolKeyOf(core.Float, []int{8})
if !ok {
t.Fatal("poolKeyOf refused a float shape of 8")
}
keep, err := core.Zeros(core.Float, 8)
if err != nil {
t.Fatal(err)
}
// Bucket cap: a full bucket refuses the next release even when the
// pool's element budget has room.
gradPool.Lock()
savedB, savedE := gradPool.buckets[k], gradPool.elems
gradPool.buckets[k] = make([]*core.Array, gradPoolPerBucket)
gradPool.elems = gradPoolMaxElems - 16
gradPool.Unlock()
releaseGrad(keep, []int{8})
gradPool.Lock()
gotLen := len(gradPool.buckets[k])
gradPool.buckets[k], gradPool.elems = savedB, savedE
gradPool.Unlock()
if gotLen != gradPoolPerBucket {
t.Fatalf("bucket accepted a release past its cap: %d entries, cap %d", gotLen, gradPoolPerBucket)
}
// Element cap: a full pool refuses the next release however empty
// the bucket is.
gradPool.Lock()
gradPool.buckets[k] = nil
gradPool.elems = gradPoolMaxElems
gradPool.Unlock()
releaseGrad(keep, []int{8})
gradPool.Lock()
gotElems := gradPool.elems
gradPool.buckets[k], gradPool.elems = savedB, savedE
gradPool.Unlock()
if gotElems != gradPoolMaxElems {
t.Fatalf("pool accepted a release past its element cap: %d, cap %d", gotElems, gradPoolMaxElems)
}
}
+80
View File
@@ -0,0 +1,80 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import "sourcedock.dev/petrbalvin/tensor/internal/core"
// Slice extracts a range along the given dimension as a new tensor;
// the backward writes the incoming gradient into the corresponding
// region of the original shape.
func (t *Tensor) Slice(dim, start, stop int) (*Tensor, error) {
if err := t.checkDiff("Slice"); err != nil {
return nil, err
}
out, err := core.Slice(t.data, dim, start, stop)
if err != nil {
return nil, err
}
orig := t.data.Shape()
dt := t.data.Dtype()
return t.unaryResult("Slice", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
// The narrowing Concat applies: a complex gradient reaching a
// real slice narrows by 2·Re before the span is copied. A
// slice's output dtype equals its input's, so only a complex
// gradient on a real tensor can differ here.
gn, err := narrowGradient(g, dt)
if err != nil {
return err
}
da := gradSlot{arr: ar.borrowGrad(dt, orig), sh: orig}
outer := 1
for d := range dim {
outer *= orig[d]
}
inner := 1
for d := dim + 1; d < len(orig); d++ {
inner *= orig[d]
}
nS := stop - start
// Each kept row is one contiguous inner run, so matching dtypes
// ride raw slice moves instead of per-element accessor calls.
fast := !gn.arr.Strided() && gn.arr.Dtype() == dt && dt != core.Int
for o := range outer {
for si := range nS {
d := o*orig[dim]*inner + (start+si)*inner
s := o*nS*inner + si*inner
if fast {
copySegRaw(da.arr, gn.arr, d, s, inner)
continue
}
for j := range inner {
copyElem(da.arr, d+j, gn.arr, s+j)
}
}
}
dst[0] = da
return nil
}), nil
}
// Reshape returns a new view-equivalent tensor of the given shape; the
// backward simply reshapes the incoming gradient back.
func (t *Tensor) Reshape(shape ...int) (*Tensor, error) {
if err := t.checkDiff("Reshape"); err != nil {
return nil, err
}
out, err := core.Reshape(t.data, shape...)
if err != nil {
return nil, err
}
orig := append([]int{}, t.data.Shape()...)
return t.unaryResult("Reshape", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
gr, err := core.Reshape(g.arr, orig...)
if err != nil {
return err
}
dst[0] = gradSlot{arr: gr, sh: orig}
return nil
}), nil
}
+246
View File
@@ -0,0 +1,246 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
"sourcedock.dev/petrbalvin/tensor/signal"
)
// Spectral autograd: the Fourier transforms as graph nodes,
// so deconvolution, spectral de-noising and frequency-domain fitting
// differentiate end to end. The adjoint of the unnormalised forward
// DFT y = F·z is dz = Fᴴ·g = n·IFFT(g) in the Wirtinger convention
// (the conjugate transpose falls out of dz = 2Re[ḡᵀ·dy] exactly the
// way the MatMul adjoint does); the inverse transform is its mirror.
// Real inputs flow through unchanged: signal.FFT widens them to
// complex, and the engine's complex-to-real narrowing (2·Re) is
// precisely the adjoint of that widening.
// FFT is the forward discrete Fourier transform of a rank-1 tensor;
// the backward multiplies the incoming gradient by Fᴴ, which is the
// inverse transform scaled by n.
func (t *Tensor) FFT() (*Tensor, error) {
if err := t.checkDiff("FFT"); err != nil {
return nil, err
}
if t.data.NDim() != 1 {
return nil, errf("autograd FFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.FFT(t.data)
if err != nil {
return nil, err
}
scale := complex(float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("FFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
inv, err := signal.IFFT(g.arr)
if err != nil {
return err
}
cs := inv.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: inv, sh: sh}
return nil
}), nil
}
// IFFT is the inverse transform of a rank-1 tensor; the backward runs
// the forward transform scaled by 1/n.
func (t *Tensor) IFFT() (*Tensor, error) {
if err := t.checkDiff("IFFT"); err != nil {
return nil, err
}
if t.data.NDim() != 1 {
return nil, errf("autograd IFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.IFFT(t.data)
if err != nil {
return nil, err
}
scale := complex(1/float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("IFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
fwd, err := signal.FFT(g.arr)
if err != nil {
return err
}
cs := fwd.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: fwd, sh: sh}
return nil
}), nil
}
// FFT2 is the 2-D forward transform; the backward is the 2-D inverse
// scaled by H·W, the total element count.
func (t *Tensor) FFT2() (*Tensor, error) {
if err := t.checkDiff("FFT2"); err != nil {
return nil, err
}
if t.data.NDim() != 2 {
return nil, errf("autograd FFT2: needs a rank-2 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.FFT2(t.data)
if err != nil {
return nil, err
}
scale := complex(float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("FFT2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
inv, err := signal.IFFT2(g.arr)
if err != nil {
return err
}
cs := inv.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: inv, sh: sh}
return nil
}), nil
}
// IFFT2 is the 2-D inverse transform; the backward is the 2-D forward
// scaled by 1/(H·W).
func (t *Tensor) IFFT2() (*Tensor, error) {
if err := t.checkDiff("IFFT2"); err != nil {
return nil, err
}
if t.data.NDim() != 2 {
return nil, errf("autograd IFFT2: needs a rank-2 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.IFFT2(t.data)
if err != nil {
return nil, err
}
scale := complex(1/float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("IFFT2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
fwd, err := signal.FFT2(g.arr)
if err != nil {
return err
}
cs := fwd.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: fwd, sh: sh}
return nil
}), nil
}
// RFFT is the real-input half-spectrum transform. The input must be a
// rank-1 real tensor; the backward folds the incoming half-spectrum
// gradient into dx = 2·Re(F_halfᴴ·g), evaluated by one padded forward
// FFT so the cost matches the forward transform. The factor 2 lands
// only on the mirrored bins through the zero padding, exactly the
// combinatorics the derivation gives.
func (t *Tensor) RFFT() (*Tensor, error) {
if err := t.checkDiff("RFFT"); err != nil {
return nil, err
}
if isComplexArr(t.data) {
return nil, errf("autograd RFFT: needs a real tensor, got %s", t.data.Dtype())
}
if t.data.NDim() != 1 {
return nil, errf("autograd RFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.RFFT(t.data)
if err != nil {
return nil, err
}
in := t.data
return t.unaryResult("RFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
n := in.Len()
half := n/2 + 1
// The adjoint needs Σ_{k<half} g_k·e^{+2πijk/n}, a +sign DFT
// of the zero-padded gradient: conj(FFT(conj(·))).
pad := make([]complex128, n)
for k := range half {
pad[k] = conj(g.arr.ComplexAt(k))
}
padArr, err := core.ComplexFromArray(pad, n)
if err != nil {
return err
}
spec, err := signal.FFT(padArr)
if err != nil {
return err
}
sh := in.Shape()
dx := gradSlot{arr: ar.borrowGrad(in.Dtype(), sh), sh: sh}
// The transform's output and dx are both freshly allocated and
// dense, so the doubled real part is taken from the payload.
ss := spec.RawComplexes()
if in.Dtype() == core.Float32 {
ds := dx.arr.RawFloat32s()[:dx.arr.Len()]
for j := range n {
ds[j] = float32(2 * real(ss[j]))
}
} else {
ds := dx.arr.RawFloats()[:dx.arr.Len()]
for j := range n {
ds[j] = 2 * real(ss[j])
}
}
dst[0] = dx
return nil
}), nil
}
// IRFFT is the inverse half-spectrum transform: a rank-1 complex
// tensor of n/2+1 bins into a real signal of length n. The backward
// widens the real gradient to a full forward FFT and halves the
// self-mirrored bins (DC, and Nyquist when n is even): dIn_k =
// FFT(g)_k/n on the ordinary bins and half of that on the mirrored
// bins, the transpose of the Hermitian extension the forward performs.
func (t *Tensor) IRFFT(n int) (*Tensor, error) {
if !isComplexArr(t.data) {
return nil, errf("autograd IRFFT: needs a complex tensor, got %s", t.data.Dtype())
}
if t.data.NDim() != 1 {
return nil, errf("autograd IRFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.IRFFT(t.data, n)
if err != nil {
return nil, err
}
in := t.data
return t.unaryResult("IRFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
gs := make([]complex128, n)
for j := range n {
gs[j] = complex(g.arr.FloatAt(j), 0)
}
gArr, err := core.ComplexFromArray(gs, n)
if err != nil {
return err
}
spec, err := signal.FFT(gArr)
if err != nil {
return err
}
half := in.Len()
full := complex(float64(n), 0)
sh := []int{half}
dx := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
ds := dx.arr.RawComplexes()[:half]
ss := spec.RawComplexes()
ds[0] = ss[0] / (2 * full)
for k := 1; k < half; k++ {
if n%2 == 0 && k == half-1 {
// Nyquist mirrors itself.
ds[k] = ss[k] / (2 * full)
continue
}
ds[k] = ss[k] / full
}
dst[0] = dx
return nil
}), nil
}
+370
View File
@@ -0,0 +1,370 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// spectralWeights builds a deterministic complex weight vector used to
// fold a spectrum into a real scalar loss.
func spectralWeights(n int, seed int) []complex128 {
g := core.NewGenerator(int64(seed))
w := make([]complex128, n)
for i := range n {
w[i] = complex(g.NormalUnit(), g.NormalUnit())
}
return w
}
// complexLossOf runs the op chain and folds the result into the real
// scalar Σ Re(w·y) + Σ|y|²/len, the same fold the graph's foldReal
// builds from Mul and Real, so oracle and graph define one loss.
func complexLossOf(t *testing.T, op func(*Tensor) (*Tensor, error), w []complex128) func(*core.Array) float64 {
t.Helper()
return func(a *core.Array) float64 {
y, err := op(FromArray(a, false))
if err != nil {
t.Fatalf("forward: %v", err)
}
s := 0.0
for i := range y.Data().Len() {
z := y.Data().ComplexAt(i)
s += real(w[i%len(w)] * z)
s += (real(z)*real(z) + imag(z)*imag(z)) / float64(y.Data().Len())
}
return s
}
}
// realLossOf is complexLossOf for chains that end in a real tensor.
func realLossOf(t *testing.T, op func(*Tensor) (*Tensor, error), w []complex128) func(*core.Array) float64 {
t.Helper()
return func(a *core.Array) float64 {
y, err := op(FromArray(a, false))
if err != nil {
t.Fatalf("forward: %v", err)
}
s := 0.0
for i := range y.Data().Len() {
v := y.Data().FloatAt(i)
s += real(w[i%len(w)])*v + v*v/float64(y.Data().Len())
}
return s
}
}
// backwardComplex runs op on a fresh tensor over vals and returns the
// leaf after Backward.
func backwardComplex(t *testing.T, op func(*Tensor) (*Tensor, error), vals []complex128, shape ...int) *Tensor {
t.Helper()
a, err := core.FromComplexes(vals, shape...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
xt := FromArray(a, true)
y, err := op(xt)
if err != nil {
t.Fatalf("forward: %v", err)
}
// Fold to a real scalar so Backward has its seed.
w := spectralWeights(y.Data().Len(), 11)
var loss *Tensor
loss, err = foldReal(t, y, w)
if err != nil {
t.Fatalf("fold: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
return xt
}
// foldReal reduces a complex tensor to Σ Re(w̄·y) + Σ|y|²/len through
// the graph ops, and a real tensor to the analogous real fold.
func foldReal(t *testing.T, y *Tensor, w []complex128) (*Tensor, error) {
t.Helper()
n := y.Data().Len()
wa, err := core.FromComplexes(w, y.Data().Shape()...)
if err != nil {
return nil, err
}
wt := FromArray(wa, false)
if isComplexArr(y.Data()) {
prod, err := y.Mul(wt)
if err != nil {
return nil, err
}
re, err := prod.Real()
if err != nil {
return nil, err
}
s1, err := re.Sum()
if err != nil {
return nil, err
}
sq, err := y.Abs2()
if err != nil {
return nil, err
}
s2, err := sq.Sum()
if err != nil {
return nil, err
}
s2s, err := s2.Scale(1 / float64(n))
if err != nil {
return nil, err
}
return s1.Add(s2s)
}
prod, err := y.Mul(wt)
if err != nil {
return nil, err
}
re, err := prod.Real()
if err != nil {
return nil, err
}
s1, err := re.Sum()
if err != nil {
return nil, err
}
sq, err := y.Abs2()
if err != nil {
return nil, err
}
s2, err := sq.Sum()
if err != nil {
return nil, err
}
s2s, err := s2.Scale(1 / float64(n))
if err != nil {
return nil, err
}
return s1.Add(s2s)
}
// TestGradFFTWirtinger pins the FFT adjoint against central
// differences on a power-of-two and a Bluestein length.
func TestGradFFTWirtinger(t *testing.T) {
for _, n := range []int{8, 12} {
g := core.NewGenerator(int64(n))
vals := make([]complex128, n)
for i := range n {
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
}
op := func(x *Tensor) (*Tensor, error) { return x.FFT() }
xt := backwardComplex(t, op, vals, n)
loss := complexLossOf(t, op, spectralWeights(n, 11))
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
}
}
// TestGradIFFTWirtinger pins the IFFT adjoint.
func TestGradIFFTWirtinger(t *testing.T) {
n := 10
g := core.NewGenerator(3)
vals := make([]complex128, n)
for i := range n {
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
}
op := func(x *Tensor) (*Tensor, error) { return x.IFFT() }
xt := backwardComplex(t, op, vals, n)
loss := complexLossOf(t, op, spectralWeights(n, 11))
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
}
// TestGradFFT2Wirtinger pins the 2-D adjoint.
func TestGradFFT2Wirtinger(t *testing.T) {
rows, cols := 3, 4
g := core.NewGenerator(5)
vals := make([]complex128, rows*cols)
for i := range vals {
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
}
op := func(x *Tensor) (*Tensor, error) { return x.FFT2() }
xt := backwardComplex(t, op, vals, rows, cols)
loss := complexLossOf(t, op, spectralWeights(rows*cols, 11))
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
}
// TestGradFFTRealInput pins the 2·Re narrowing path: a real leaf under
// a complex FFT node.
func TestGradFFTRealInput(t *testing.T) {
n := 8
g := core.NewGenerator(9)
vals := make([]float64, n)
for i := range n {
vals[i] = g.NormalUnit()
}
op := func(x *Tensor) (*Tensor, error) { return x.FFT() }
a, err := core.FromFloats(vals, n)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(a, true)
y, err := op(xt)
if err != nil {
t.Fatalf("forward: %v", err)
}
loss, err := foldReal(t, y, spectralWeights(n, 11))
if err != nil {
t.Fatalf("fold: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := realToComplexLoss(t, op, n)
want := numericGrad(lossOf, a)
for i := range n {
if math.Abs(xt.Grad().FloatAt(i)-want[i]) > 1e-6 {
t.Fatalf("grad[%d] = %g, want %g", i, xt.Grad().FloatAt(i), want[i])
}
}
}
// realToComplexLoss adapts a chain over real input for numericGrad: it
// re-runs the forward and folds the (complex) output into a scalar.
func realToComplexLoss(t *testing.T, op func(*Tensor) (*Tensor, error), n int) func(*core.Array) float64 {
t.Helper()
w := spectralWeights(n, 11)
return func(a *core.Array) float64 {
y, err := op(FromArray(a, false))
if err != nil {
t.Fatalf("forward: %v", err)
}
s := 0.0
for i := range n {
z := y.Data().ComplexAt(i)
s += real(w[i] * z)
s += (real(z)*real(z) + imag(z)*imag(z)) / float64(n)
}
return s
}
}
// TestGradRFFT pins the half-spectrum adjoint, even and odd lengths.
func TestGradRFFT(t *testing.T) {
for _, n := range []int{8, 9} {
g := core.NewGenerator(int64(n * 2))
vals := make([]float64, n)
for i := range n {
vals[i] = g.NormalUnit()
}
a, err := core.FromFloats(vals, n)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(a, true)
y, err := xt.RFFT()
if err != nil {
t.Fatalf("RFFT: %v", err)
}
w := spectralWeights(y.Data().Len(), 11)
var loss *Tensor
loss, err = foldReal(t, y, w)
if err != nil {
t.Fatalf("fold: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
yp, err := FromArray(a, false).RFFT()
if err != nil {
t.Fatalf("RFFT: %v", err)
}
s := 0.0
for i := range yp.Data().Len() {
z := yp.Data().ComplexAt(i)
s += real(w[i] * z)
s += (real(z)*real(z) + imag(z)*imag(z)) / float64(yp.Data().Len())
}
return s
}
want := numericGrad(lossOf, a)
for i := range n {
if math.Abs(xt.Grad().FloatAt(i)-want[i]) > 1e-6 {
t.Fatalf("n=%d grad[%d] = %g, want %g", n, i, xt.Grad().FloatAt(i), want[i])
}
}
}
}
// TestGradIRFFT pins the half-spectrum inverse adjoint, on an even
// length (Nyquist half-weight bin) and an odd one (the last bin is an
// ordinary mirrored bin with the full 1/n weight).
func TestGradIRFFT(t *testing.T) {
for _, n := range []int{8, 9} {
half := n/2 + 1
g := core.NewGenerator(21)
vals := make([]complex128, half)
for i := range half {
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
}
op := func(x *Tensor) (*Tensor, error) { return x.IRFFT(n) }
xt := backwardComplex(t, op, vals, half)
// numericGrad over complex perturbations, folded through the real
// output.
w := spectralWeights(n, 11)
loss := func(a *core.Array) float64 {
y, err := FromArray(a, false).IRFFT(n)
if err != nil {
t.Fatalf("IRFFT: %v", err)
}
s := 0.0
for i := range n {
v := y.Data().FloatAt(i)
s += real(w[i])*v + v*v/float64(n)
}
return s
}
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
}
}
// TestGradFFTRoundtripIdentity pins the composition: gradient through
// IFFT∘FFT must arrive unchanged (Fᴴ·(1/n)F = I).
func TestGradFFTRoundtripIdentity(t *testing.T) {
n := 8
vals := make([]complex128, n)
g := core.NewGenerator(4)
for i := range n {
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
}
a, err := core.FromComplexes(vals, n)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
xt := FromArray(a, true)
f, err := xt.FFT()
if err != nil {
t.Fatalf("FFT: %v", err)
}
fi, err := f.IFFT()
if err != nil {
t.Fatalf("IFFT: %v", err)
}
w := spectralWeights(n, 7)
loss, err := foldReal(t, fi, w)
if err != nil {
t.Fatalf("fold: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// dL/dy at the roundtrip output, propagated through both adjoints,
// must equal dL/dy itself: the fold's Wirtinger gradient is
// w̄/2 + y/n (Mul+Real contributes w̄/2, Abs2/n contributes y/n).
for i := range n {
dy := conj(w[i])/2 + fi.Data().ComplexAt(i)/complex(float64(n), 0)
got := xt.Grad().ComplexAt(i)
if cmplxAbs(got-dy) > 1e-8 {
t.Fatalf("grad[%d] = %v, want %v", i, got, dy)
}
}
}
+286
View File
@@ -0,0 +1,286 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The cross-rank gradient sweep (the machine that catches regressions
// hiding in untested shapes, like the LayerNorm affine reduction):
// every listed differentiable op runs through a finite-difference
// check on several input ranks and both float element types.
type sweepCase struct {
name string
ranks [][]int // the shapes the op must answer for
run func(x *Tensor) (*Tensor, error)
}
// sweepValue is deterministic, sign-varying and comfortably away from
// kinks and saturation boundaries.
func sweepValue(i int) float64 {
return math.Sin(float64(i%17)*0.7)*2 + 0.25
}
func sweepPattern(n int) []float64 {
out := make([]float64, n)
for i := range out {
out[i] = 0.5*float64(i%4) - 0.75
}
return out
}
func sweepMatrixPattern(width int) *core.Array {
vals := sweepPattern(width * width)
arr, _ := core.FromFloats(vals, width, width)
return arr
}
// sweepPositive keeps logs and divisions inside their real domains no
// matter how the signed sweep values land, shaped like the input.
func sweepPositive(shape []int) *core.Array {
n := numEl(shape)
out := make([]float64, n)
for i := range out {
out[i] = 3 + float64(i%3)
}
arr, _ := core.FromFloats(out, shape...)
return arr
}
// sweepSecond derives a second operand from an independent pattern:
// paired cases need two leaves but stay deterministic.
func sweepSecond(shape []int) (*Tensor, error) {
total := 1
for _, d := range shape {
total *= d
}
vals := make([]float64, total)
for i := range vals {
vals[i] = math.Cos(float64(i%11))*1.5 - 0.5
}
a, err := core.FromFloats(vals, shape...)
if err != nil {
return nil, err
}
return FromArray(a, true), nil
}
func numEl(shape []int) int {
n := 1
for _, d := range shape {
n *= d
}
return n
}
// checkOp runs one case per shape and dtype: forward on a gradient
// leaf, weighted-sum loss with a fixed mask so every slot gets its own
// coefficient, then analytic-vs-central-difference compare.
func checkOp(t *testing.T, tc sweepCase, dt core.Dtype) {
t.Helper()
for _, dims := range tc.ranks {
n := numEl(dims)
build := func(reqGrad bool) *Tensor {
vals := make([]float64, n)
for i := range vals {
vals[i] = sweepValue(i + len(dims))
}
a, err := core.FromFloats(vals, dims...)
if err != nil {
t.Fatal(err)
}
if dt == core.Float32 {
a32, cerr := core.Astype(a, core.Float32)
if cerr != nil {
t.Fatal(cerr)
}
a = a32
}
return FromArray(a, reqGrad)
}
x := build(true)
out, err := tc.run(x)
if err != nil {
t.Fatalf("%s %v %v: forward: %v", tc.name, dims, dt, err)
}
// The mask matches the OUTPUT shape: reducing ops return fewer
// slots than their input carries.
outN := out.Data().Len()
maskVals := sweepPattern(outN)
mArr, _ := core.FromFloats(maskVals, out.Data().Shape()...)
scaled, err := out.Mul(FromArray(mArr, false))
if err != nil {
t.Fatalf("%s %v %v: loss mul: %v", tc.name, dims, dt, err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatalf("%s %v %v: loss sum: %v", tc.name, dims, dt, err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("%s %v %v: backward: %v", tc.name, dims, dt, err)
}
ref := numericGrad(func(v *core.Array) float64 {
o, rerr := tc.run(FromArray(v, false))
if rerr != nil {
return math.NaN()
}
total := 0.0
for i := range maskVals {
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
}
return total
}, x.Data())
got := x.Grad()
if got.Len() != len(ref) {
t.Fatalf("%s %v %v: gradient length %d, reference %d",
tc.name, dims, dt, got.Len(), len(ref))
}
scale := 1.0
for _, r := range ref {
if s := math.Abs(r); s > scale {
scale = s
}
}
tol := 1e-4
if dt == core.Float32 {
tol = 8e-2
}
for i := range ref {
if math.Abs(got.FloatAt(i)-ref[i]) > tol*scale {
t.Errorf("%s %v %v: grad[%d] = %v, want ≈%v",
tc.name, dims, dt, i, got.FloatAt(i), ref[i])
}
}
}
}
func TestGradientSweepAcrossRanksAndDtypes(t *testing.T) {
shapes234 := [][]int{{4}, {2, 3}, {2, 2, 2}}
simple := []sweepCase{
{name: "Neg", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Neg() }},
{name: "Exp", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Exp() }},
{name: "Sigmoid", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Sigmoid() }},
{name: "Tanh", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Tanh() }},
{name: "Abs", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Abs() }},
{name: "Pow3", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Pow(3) }},
{name: "Scale", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Scale(-1.75) }},
{name: "ClipInterior", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Clip(-2.75, 2.75) }},
{name: "LogShifted", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) {
up, err := x.Add(FromArray(sweepPositive(x.Data().Shape()), false))
if err != nil {
return nil, err
}
return up.Log()
}},
{name: "TransposeAxesReverse", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
d := x.Data().NDim()
perm := make([]int, d)
for i := range perm {
perm[i] = d - 1 - i
}
return x.TransposeAxes(perm...)
}},
{name: "ReshapeFlatten", ranks: [][]int{{2, 3}, {2, 2, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.Reshape(x.Data().Len())
}},
{name: "SumAxisZero", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.SumAxis(0)
}},
{name: "MeanAxisLast", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.MeanAxis(x.Data().NDim() - 1)
}},
// The L2 norm's backward runs a dedicated float32 sweep beside
// the float64 one; the two dtype legs below drive both.
{name: "L2NormAxisLast", ranks: [][]int{{4}, {2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.L2NormAxis(x.Data().NDim() - 1)
}},
}
elementPairs := []struct {
name string
op func(a, b *Tensor) (*Tensor, error)
}{
{"Add", func(a, b *Tensor) (*Tensor, error) { return a.Add(b) }},
{"Sub", func(a, b *Tensor) (*Tensor, error) { return a.Sub(b) }},
{"Mul", func(a, b *Tensor) (*Tensor, error) { return a.Mul(b) }},
}
for _, ep := range elementPairs {
simple = append(simple, sweepCase{
name: ep.name,
ranks: shapes234,
run: func(x *Tensor) (*Tensor, error) {
other, err := sweepSecond(x.Data().Shape())
if err != nil {
return nil, err
}
return ep.op(x, other)
},
})
}
// Division keeps both operands positive via the shared shift.
simple = append(simple, sweepCase{
name: "DivShifted",
ranks: shapes234,
run: func(x *Tensor) (*Tensor, error) {
other, err := sweepSecond(x.Data().Shape())
if err != nil {
return nil, err
}
lift, err := other.Add(FromArray(sweepPositive(x.Data().Shape()), false))
if err != nil {
return nil, err
}
return x.Div(lift)
},
})
// Column concatenation against half of a second leaf.
simple = append(simple, sweepCase{
name: "ConcatColumns",
ranks: [][]int{{2, 4}},
run: func(x *Tensor) (*Tensor, error) {
extra, err := sweepSecond(x.Data().Shape())
if err != nil {
return nil, err
}
halves, err := extra.Slice(1, 0, x.Data().Shape()[1]/2)
if err != nil {
return nil, err
}
return x.Concat(halves, 1)
},
})
// A matmul product collapsed by an axis sum, the inference
// backbone's gradient path.
simple = append(simple, sweepCase{
name: "MatMulSumRows",
ranks: [][]int{{3, 4}},
run: func(x *Tensor) (*Tensor, error) {
cols := x.Data().Shape()[1]
product, err := x.MatMul(FromArray(sweepMatrixPattern(cols), false))
if err != nil {
return nil, err
}
return product.SumAxis(0)
},
})
for _, tc := range simple {
for _, dt := range []core.Dtype{core.Float, core.Float32} {
checkOp(t, tc, dt)
}
}
}
+2303
View File
File diff suppressed because it is too large Load Diff
+343
View File
@@ -0,0 +1,343 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"strings"
"testing"
)
func mustTensor(t *testing.T, vals []float64, shape ...int) *Tensor {
t.Helper()
tt, err := FromFloat64s(vals, true, shape...)
if err != nil {
t.Fatalf("FromFloat64s(%v, %v): %v", vals, shape, err)
}
return tt
}
func TestAutogradBasicChain(t *testing.T) {
x := mustTensor(t, []float64{2, 3}, 2)
y := mustTensor(t, []float64{4, 5}, 2)
z, err := x.Add(y)
if err != nil {
t.Fatal(err)
}
sq, err := z.Mul(z)
if err != nil {
t.Fatal(err)
}
s, err := sq.Sum()
if err != nil {
t.Fatal(err)
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
// d/dx (x+y)^2 summed = 2(x+y); at x=2: 12, at x=3: 16.
if gx := x.Grad().FloatAt(0); math.Abs(gx-12) > 1e-9 {
t.Errorf("grad x[0]: got %v, want 12", gx)
}
if gx := x.Grad().FloatAt(1); math.Abs(gx-16) > 1e-9 {
t.Errorf("grad x[1]: got %v, want 16", gx)
}
// y's gradient matches x's, symmetric in the sum.
if gy := y.Grad().FloatAt(0); math.Abs(gy-12) > 1e-9 {
t.Errorf("grad y[0]: got %v, want 12", gy)
}
}
func TestAutogradMatMul(t *testing.T) {
a := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
b := mustTensor(t, []float64{5, 6, 7, 8}, 2, 2)
p, err := a.MatMul(b)
if err != nil {
t.Fatal(err)
}
s, err := p.Sum()
if err != nil {
t.Fatal(err)
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
// d/dA sum(A·B) = J·Bᵀ; Bᵀ = [[5,7],[6,8]], so J·Bᵀ =
// [[11,15],[11,15]] (each row is the column sums of Bᵀ).
wantA := []float64{11, 15, 11, 15}
for i := range 4 {
if g := a.Grad().FloatAt(i); math.Abs(g-wantA[i]) > 1e-9 {
t.Errorf("grad A[%d]: got %v, want %v", i, g, wantA[i])
}
}
// d/dB sum(A·B) = Aᵀ·J; Aᵀ = [[1,3],[2,4]], row sums: 4, 6, so
// Aᵀ·J = [[4,4],[6,6]].
wantB := []float64{4, 4, 6, 6}
for i := range 4 {
if g := b.Grad().FloatAt(i); math.Abs(g-wantB[i]) > 1e-9 {
t.Errorf("grad B[%d]: got %v, want %v", i, g, wantB[i])
}
}
}
func TestAutogradMatMulVector(t *testing.T) {
// 2-D × 1-D: y = A·x.
a := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
x := mustTensor(t, []float64{2, 3}, 2)
y, err := a.MatMul(x)
if err != nil {
t.Fatal(err)
}
s, err := y.Sum()
if err != nil {
t.Fatal(err)
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
// grad x = Aᵀ·1 = column sums of A: 4, 6.
if g := x.Grad().FloatAt(0); math.Abs(g-4) > 1e-9 {
t.Errorf("grad x[0]: got %v, want 4", g)
}
if g := x.Grad().FloatAt(1); math.Abs(g-6) > 1e-9 {
t.Errorf("grad x[1]: got %v, want 6", g)
}
// grad A = outer(1, x): [[2,3],[2,3]].
if g := a.Grad().FloatAt(2); math.Abs(g-2) > 1e-9 {
t.Errorf("grad A[2]: got %v, want 2", g)
}
}
func TestAutogradActivations(t *testing.T) {
// Sigmoid at 0: σ(0)=0.5, σ' = 0.25.
x3 := mustTensor(t, []float64{0}, 1)
sg, _ := x3.Sigmoid()
if err := sg.Backward(); err != nil {
t.Fatal(err)
}
if g := x3.Grad().FloatAt(0); math.Abs(g-0.25) > 1e-9 {
t.Errorf("Sigmoid grad at 0: got %v, want 0.25", g)
}
// Exp and Log compose to identity: grad log(exp(x)) = 1.
x4 := mustTensor(t, []float64{2}, 1)
e, _ := x4.Exp()
l, _ := e.Log()
if err := l.Backward(); err != nil {
t.Fatal(err)
}
if g := x4.Grad().FloatAt(0); math.Abs(g-1) > 1e-9 {
t.Errorf("log(exp) grad: got %v, want 1", g)
}
}
func TestAutogradGradientAccumulation(t *testing.T) {
x := mustTensor(t, []float64{1}, 1)
a, _ := x.Mul(x)
b, _ := x.Mul(x)
s, err := a.Add(b)
if err != nil {
t.Fatal(err)
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
// d/dx (x² + x²) at x=1 = 4.
if g := x.Grad().FloatAt(0); math.Abs(g-4) > 1e-9 {
t.Errorf("accumulated grad: got %v, want 4", g)
}
// A second Backward without ZeroGrad accumulates into the leaf.
if err := s.Backward(); err != nil {
t.Fatal(err)
}
if g := x.Grad().FloatAt(0); math.Abs(g-8) > 1e-9 {
t.Errorf("accumulated grad after 2nd pass: got %v, want 8", g)
}
x.ZeroGrad()
if x.Grad() != nil {
t.Error("ZeroGrad did not clear the gradient")
}
}
func TestAutogradRejectsNonFloat(t *testing.T) {
i, err := core.FromInts([]int64{1, 2}, 2)
if err != nil {
t.Fatal(err)
}
it := FromArray(i, true)
if _, err := it.Sum(); err == nil || !strings.Contains(err.Error(), "float") {
t.Errorf("int Sum: %v", err)
}
c, _ := core.FromComplexes([]complex128{1 + 2i}, 1)
ct := FromArray(c, true)
// Complex Exp is differentiable (the Wirtinger graph); the
// real-only kernels are the ones that must still refuse it.
if _, err := ct.Exp(); err != nil {
t.Errorf("complex Exp must differentiate: %v", err)
}
if _, err := ct.Log(); err == nil {
t.Error("complex Log must error")
}
if _, err := ct.Tanh(); err == nil {
t.Error("complex Tanh must error")
}
}
func TestAutogradDivTanhNeg(t *testing.T) {
// d/dx (x/y) at x=4, y=2 = 1/2.
x := mustTensor(t, []float64{4}, 1)
y := mustTensor(t, []float64{2}, 1)
q, err := x.Div(y)
if err != nil {
t.Fatal(err)
}
if err := q.Backward(); err != nil {
t.Fatal(err)
}
if g := x.Grad().FloatAt(0); math.Abs(g-0.5) > 1e-9 {
t.Errorf("Div grad x: got %v, want 0.5", g)
}
// d/dy (x/y) at y=2 = -x/y² = -1.
if g := y.Grad().FloatAt(0); math.Abs(g+1) > 1e-9 {
t.Errorf("Div grad y: got %v, want -1", g)
}
// tanh'(0) = 1.
t0 := mustTensor(t, []float64{0}, 1)
th, _ := t0.Tanh()
if err := th.Backward(); err != nil {
t.Fatal(err)
}
if g := t0.Grad().FloatAt(0); math.Abs(g-1) > 1e-9 {
t.Errorf("Tanh grad at 0: got %v, want 1", g)
}
// d/dx (-x) = -1.
n := mustTensor(t, []float64{3}, 1)
neg, _ := n.Neg()
if err := neg.Backward(); err != nil {
t.Fatal(err)
}
if g := n.Grad().FloatAt(0); g != -1 {
t.Errorf("Neg grad: got %v, want -1", g)
}
// Accessors and Mean grad.
m := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
if m.Data() != m.Data() || m.RequiresGrad() != true {
t.Error("accessors wrong")
}
mean, err := m.Mean()
if err != nil {
t.Fatal(err)
}
if err := mean.Backward(); err != nil {
t.Fatal(err)
}
// d/dx mean(x) = 1/n = 1/4.
if g := m.Grad().FloatAt(0); math.Abs(g-0.25) > 1e-9 {
t.Errorf("Mean grad: got %v, want 0.25", g)
}
if m.Grad() == nil {
t.Error("Grad() must be non-nil after Backward")
}
}
func TestAutogradReshape(t *testing.T) {
x, err := FromFloat64s([]float64{1, 2, 3, 4}, true, 2, 2)
if err != nil {
t.Fatal(err)
}
r, err := x.Reshape(4)
if err != nil {
t.Fatal(err)
}
sumT, err := r.Sum()
if err != nil {
t.Fatal(err)
}
if err := sumT.Backward(); err != nil {
t.Fatal(err)
}
g, err := x.Grad().Elements[float64]()
if err != nil {
t.Fatal(err)
}
for i := range g {
if g[i] != 1 {
t.Fatalf("Reshape grad[%d]: %v, want 1", i, g[i])
}
}
}
func TestAutogradTransposeBackward(t *testing.T) {
x, err := FromFloat64s([]float64{1, 2, 3, 4}, true, 2, 2)
if err != nil {
t.Fatal(err)
}
tr, err := x.Transpose()
if err != nil {
t.Fatal(err)
}
s, err := tr.Sum()
if err != nil {
t.Fatal(err)
}
if err := s.Backward(); err != nil {
t.Fatal(err)
}
g, err := x.Grad().Elements[float64]()
if err != nil {
t.Fatal(err)
}
for i := range g {
if g[i] != 1 {
t.Errorf("Transpose grad[%d]: %v, want 1", i, g[i])
}
}
}
// TestAutogradPowZeroGradient pins the exponent-0 backward: d/dx x⁰
// is the zero gradient everywhere, including at x = 0 where the
// chain rule would evaluate 0·∞ and produce NaN.
func TestAutogradPowZeroGradient(t *testing.T) {
x := mustTensor(t, []float64{0, 2}, 2)
y, err := x.Pow(0)
if err != nil {
t.Fatalf("Pow(0): %v", err)
}
loss, err := y.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
for i := range 2 {
g := x.Grad().FloatAt(i)
if math.IsNaN(g) || g != 0 {
t.Errorf("d/dx x⁰ at %g = %v, want exactly 0", x.Data().FloatAt(i), g)
}
}
}
// TestAutogradLeafBackwardAccumulates pins that Backward on a leaf
// accumulates into the existing gradient like any other backward
// pass, instead of overwriting it.
func TestAutogradLeafBackwardAccumulates(t *testing.T) {
x := mustTensor(t, []float64{3}, 1)
if err := x.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
if got := x.Grad().FloatAt(0); got != 1 {
t.Fatalf("first leaf Backward: grad %v, want 1", got)
}
if err := x.Backward(); err != nil {
t.Fatalf("second Backward: %v", err)
}
if got := x.Grad().FloatAt(0); got != 2 {
t.Fatalf("second leaf Backward: grad %v, want 2 (accumulated)", got)
}
}
+256
View File
@@ -0,0 +1,256 @@
// 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"
"sourcedock.dev/petrbalvin/tensor/linalg"
)
// Benchmarks for the per-step scratch of the stiff solvers, the PDE
// stencil steps and the finite-element assemblies: the paths where
// allocation churn and repeated lookups, not the arithmetic, set the
// cost.
// perfVector wraps a fixed literal as a rank-1 array.
func perfVector(b *testing.B, vals []float64) *core.Array {
b.Helper()
a, err := core.FromFloats(vals, len(vals))
if err != nil {
b.Fatal(err)
}
return a
}
// perfStiffDecay builds the diagonal stiff system y' = −100(i+1)·y_i
// with every component started at one: the rates span three decades,
// so the step control stretches over the fast transient and the
// Jacobian stays diagonal and cheap to evaluate.
func perfStiffDecay(n int) func(t float64, y *core.Array) (*core.Array, error) {
rates := make([]float64, n)
for i := range rates {
rates[i] = -100 * float64(i+1)
}
return func(t float64, y *core.Array) (*core.Array, error) {
out := core.New(core.Float, n)
vals := out.RawFloats()
ys := y.RawFloats()
for i := range n {
vals[i] = rates[i] * ys[i]
}
return out, nil
}
}
// perfConstantState returns a vector of n ones.
func perfConstantState(n int) []float64 {
vals := make([]float64, n)
for i := range vals {
vals[i] = 1
}
return vals
}
func BenchmarkROS4Stiff(b *testing.B) {
const n = 32
f := perfStiffDecay(n)
start := perfVector(b, perfConstantState(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)
}
}
}
func BenchmarkBDFVarStiff(b *testing.B) {
const n = 32
f := perfStiffDecay(n)
start := perfVector(b, perfConstantState(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)
}
}
}
func BenchmarkHeat1DStepLoop(b *testing.B) {
const n = 256
u0 := make([]float64, n)
for i := range u0 {
u0[i] = math.Sin(float64(i+1) / float64(n+1) * math.Pi)
}
state := perfVector(b, u0)
// Two samples put every step inside the loop under test: the
// published history costs one copy either way.
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateHeat1D(state, 1, 1.0/257, 0.05, 1e-4, 2, 0, 0); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkWave1DStepLoop(b *testing.B) {
const n = 256
u0 := make([]float64, n)
v0 := make([]float64, n)
for i := range u0 {
u0[i] = math.Sin(float64(i+1) / float64(n+1) * math.Pi)
}
state := perfVector(b, u0)
vel := perfVector(b, v0)
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateWave1D(state, vel, 1, 1.0/257, 0.05, 1e-4, 2); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkHeat2DStepLoop(b *testing.B) {
const rows, cols = 64, 64
u0 := make([]float64, rows*cols)
for r := range rows {
for c := range cols {
u0[r*cols+c] = math.Sin(float64(c+1)/float64(cols+1)*math.Pi) *
math.Sin(float64(r+1)/float64(rows+1)*math.Pi)
}
}
state, err := core.FromFloats(u0, rows, cols)
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateHeat2D(state, 1, 1.0/65, 1.0/65, 0.002, 2e-5, 2, 0, 0, 0, 0); err != nil {
b.Fatal(err)
}
}
}
// perfSquareBoundary lists the boundary nodes of the m by m cell grid
// on the unit square: the bottom and top rows, then the interior
// nodes of the left and right columns.
func perfSquareBoundary(m int) []int {
nodes := make([]int, 0, 4*m)
for i := range m + 1 {
nodes = append(nodes, i, m*(m+1)+i)
}
for j := 1; j < m; j++ {
nodes = append(nodes, j*(m+1), j*(m+1)+m)
}
return nodes
}
func BenchmarkPoissonFEM2D(b *testing.B) {
const m = 48
mesh, err := GridTriangleMesh2D(0, 0, 1, 1, m, m)
if err != nil {
b.Fatal(err)
}
bound := perfSquareBoundary(m)
values := make([]float64, len(bound))
opts := FEMPoissonOptions{
Kappa: 1,
DirichletNodes: bound,
DirichletValues: values,
Ordering: linalg.SparseOrderingReverseCuthillMcKee,
}
b.ReportAllocs()
for b.Loop() {
if _, err := SolvePoissonFEM2D(mesh, nil, opts); err != nil {
b.Fatal(err)
}
}
}
// perfBoxBoundary lists the vertices of the box tetrahedral mesh that
// sit on the unit cube's surface.
func perfBoxBoundary(mesh *TetraMesh3D) []int {
nodes := make([]int, 0, mesh.Vertices3())
for i := range mesh.Vertices3() {
x, y, z := mesh.Vertices[3*i], mesh.Vertices[3*i+1], mesh.Vertices[3*i+2]
if x == 0 || x == 1 || y == 0 || y == 1 || z == 0 || z == 1 {
nodes = append(nodes, i)
}
}
return nodes
}
func BenchmarkPoissonFEM3D(b *testing.B) {
const m = 8
mesh, err := BoxTetraMesh3D(0, 0, 0, 1, 1, 1, m, m, m)
if err != nil {
b.Fatal(err)
}
bound := perfBoxBoundary(mesh)
values := make([]float64, len(bound))
opts := FEMPoisson3DOptions{
Kappa: 1,
DirichletNodes: bound,
DirichletValues: values,
Ordering: linalg.SparseOrderingReverseCuthillMcKee,
}
b.ReportAllocs()
for b.Loop() {
if _, err := SolvePoissonFEM3D(mesh, nil, opts); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkPoissonFEM3DLoad(b *testing.B) {
const m = 5
mesh, err := BoxTetraMesh3D(0, 0, 0, 1, 1, 1, m, m, m)
if err != nil {
b.Fatal(err)
}
bound := perfBoxBoundary(mesh)
values := make([]float64, len(bound))
src := func(x, y, z float64) float64 {
return 3 * math.Pi * math.Pi * math.Sin(math.Pi*x) * math.Sin(math.Pi*y) * math.Sin(math.Pi*z)
}
opts := FEMPoisson3DOptions{
Kappa: 1,
DirichletNodes: bound,
DirichletValues: values,
Ordering: linalg.SparseOrderingReverseCuthillMcKee,
}
b.ReportAllocs()
for b.Loop() {
if _, err := SolvePoissonFEM3D(mesh, src, opts); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkIntegrateHeat2DBig is the same scheme on a grid large enough
// that the step sweeps have work to share: 512 lines of 512 unknowns per
// half-step.
func BenchmarkIntegrateHeat2DBig(b *testing.B) {
const rows, cols = 512, 512
u0 := make([]float64, rows*cols)
for r := range rows {
for c := range cols {
u0[r*cols+c] = math.Sin(float64(c)/float64(cols)*math.Pi) * math.Sin(float64(r)/float64(rows)*math.Pi)
}
}
state, err := core.FromFloats(u0, rows, cols)
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateHeat2D(state, 1, 1.0/513, 1.0/513, 0.02, 0.0004, 2, 0, 0, 0, 0); err != nil {
b.Fatal(err)
}
}
}
+305
View File
@@ -0,0 +1,305 @@
// 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 package's heavy paths: the adaptive Dormand-Prince
// step loop, the stiff implicit schemes with their numerical
// Jacobians, the adaptive quadrature and cubature, and the PDE
// stencils.
// odeLinear builds the closed-form linear system y' = A·y with a
// stable diagonal A, the cheapest honest workload for an adaptive
// step loop, and returns f plus the analytic solution for callers
// that want it.
func odeLinear(n int) (func(float64, *core.Array) (*core.Array, error), []float64) {
rates := make([]float64, n)
for i := range rates {
rates[i] = -0.25 * float64(i+1)
}
f := func(t float64, y *core.Array) (*core.Array, error) {
out := core.New(core.Float, n)
vals := out.RawFloats()
ys := y.RawFloats()
for i := range n {
vals[i] = rates[i] * ys[i]
}
return out, nil
}
y0 := make([]float64, n)
for i := range y0 {
y0[i] = 1
}
return f, y0
}
func benchVector(b *testing.B, vals []float64) *core.Array {
b.Helper()
a, err := core.FromFloats(vals, len(vals))
if err != nil {
b.Fatal(err)
}
return a
}
func BenchmarkIntegrateODE(b *testing.B) {
f, y0 := odeLinear(16)
start := benchVector(b, y0)
opts := ODEOptions{}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateODE(f, 0, 10, start, opts); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateRK4(b *testing.B) {
f, y0 := odeLinear(16)
start := benchVector(b, y0)
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateRK4(f, 0, 10, start, 2000); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateBackwardEuler(b *testing.B) {
// A stiff diagonal system: rates from −1 to −1000.
const n = 4
f := func(t float64, y *core.Array) (*core.Array, error) {
out := core.New(core.Float, n)
vals := out.RawFloats()
ys := y.RawFloats()
for i := range n {
vals[i] = -float64(i+1) * 100 * ys[i]
}
return out, nil
}
start := benchVector(b, []float64{1, 1, 1, 1})
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateBackwardEuler(f, 0, 1, start, 200, ODEOptions{}); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateBDF2(b *testing.B) {
const n = 4
f := func(t float64, y *core.Array) (*core.Array, error) {
out := core.New(core.Float, n)
vals := out.RawFloats()
ys := y.RawFloats()
for i := range n {
vals[i] = -float64(i+1) * 100 * ys[i]
}
return out, nil
}
start := benchVector(b, []float64{1, 1, 1, 1})
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateBDF2(f, 0, 1, start, ODEOptions{}); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateDAE(b *testing.B) {
// The linear index-1 circuit shape: one differential row, one
// algebraic constraint, the Newton solve carrying the step.
m, err := core.FromFloats([]float64{1, 0, 0, 0}, 2, 2)
if err != nil {
b.Fatal(err)
}
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{-y.FloatAt(0), y.FloatAt(1) - y.FloatAt(0)}, 2)
}
start := benchVector(b, []float64{1, 1})
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateDAE(f, m, 0, 1, start, 200, DAEOptions{}); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateMidpoint(b *testing.B) {
// The harmonic oscillator's quadratic H: the implicit stage is a
// root find whose gradient is linear in z.
const n = 8
gradH := func(z *core.Array) (*core.Array, error) {
out := core.New(core.Float, 2*n)
vals := out.RawFloats()
zs := z.RawFloats()
for i := range n {
vals[i] = zs[n+i]
vals[n+i] = zs[i]
}
return out, nil
}
q0 := make([]float64, n)
p0 := make([]float64, n)
for i := range q0 {
q0[i] = math.Sin(float64(i))
p0[i] = math.Cos(float64(i))
}
qs := benchVector(b, q0)
ps := benchVector(b, p0)
b.ReportAllocs()
for b.Loop() {
if _, _, err := IntegrateMidpoint(gradH, 0, 1, qs, ps, 50, MidpointOptions{}); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateFunction(b *testing.B) {
f := func(x float64) (float64, error) { return math.Sin(x), nil }
opts := QuadratureOptions{}
b.ReportAllocs()
for b.Loop() {
if _, _, err := IntegrateFunction(f, 0, 100, opts); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateND(b *testing.B) {
f := func(x []float64) float64 {
s := 0.0
for _, v := range x {
s += v * v
}
return math.Exp(-s)
}
lo := []float64{-2, -2, -2}
hi := []float64{2, 2, 2}
opts := CubatureOptions{Tolerance: 1e-6}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateND(f, lo, hi, opts); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateHeat1D(b *testing.B) {
n := 256
u0 := make([]float64, n)
for i := range u0 {
u0[i] = math.Sin(float64(i) / float64(n) * math.Pi)
}
state := benchVector(b, u0)
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateHeat1D(state, 1, 1.0/257, 0.1, 0.0002, 10, 0, 0); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateWave1D(b *testing.B) {
n := 256
u0 := make([]float64, n)
v0 := make([]float64, n)
for i := range u0 {
u0[i] = math.Sin(float64(i) / float64(n) * math.Pi)
}
us := benchVector(b, u0)
vs := benchVector(b, v0)
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateWave1D(us, vs, 1, 1.0/257, 0.5, 0.002, 10); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateHeat2D(b *testing.B) {
rows, cols := 32, 32
u0 := make([]float64, rows*cols)
for r := range rows {
for c := range cols {
u0[r*cols+c] = math.Sin(float64(c)/float64(cols)*math.Pi) * math.Sin(float64(r)/float64(rows)*math.Pi)
}
}
state, err := core.FromFloats(u0, rows, cols)
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateHeat2D(state, 1, 1.0/33, 1.0/33, 0.02, 0.0004, 5, 0, 0, 0, 0); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateWave2D(b *testing.B) {
rows, cols := 32, 32
u0 := make([]float64, rows*cols)
v0 := make([]float64, rows*cols)
for r := range rows {
for c := range cols {
u0[r*cols+c] = math.Sin(float64(c)/float64(cols)*math.Pi) * math.Sin(float64(r)/float64(rows)*math.Pi)
}
}
us, err := core.FromFloats(u0, rows, cols)
if err != nil {
b.Fatal(err)
}
vs, err := core.FromFloats(v0, rows, cols)
if err != nil {
b.Fatal(err)
}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateWave2D(us, vs, 1, 1.0/33, 1.0/33, 0.05, 0.002, 5); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkIntegrateVerlet(b *testing.B) {
// Two coupled oscillators apiece: the acceleration reads the
// neighbour spring terms.
n := 32
q0 := make([]float64, n)
p0 := make([]float64, n)
for i := range q0 {
q0[i] = math.Sin(float64(i))
}
accel := func(q *core.Array) (*core.Array, error) {
out := core.New(core.Float, n)
vals := out.RawFloats()
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 := benchVector(b, q0)
ps := benchVector(b, p0)
b.ReportAllocs()
for b.Loop() {
if _, _, err := IntegrateVerlet(accel, 0, 10, qs, ps, 500); err != nil {
b.Fatal(err)
}
}
}
+75
View File
@@ -0,0 +1,75 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression pins: budgets that did not bound what they
// promised, and non-finite states that integrated to no error.
// TestIntegrateNDDimensionBudget: in 10 dimensions the root box alone
// costs 5^10 + 3^10 evaluations, about five times the default budget,
// before the first budget check could fire.
func TestIntegrateNDDimensionBudget(t *testing.T) {
lower := make([]float64, 10)
upper := make([]float64, 10)
for i := range upper {
upper[i] = 1
}
f := func(x []float64) float64 { return 1 }
_, err := IntegrateND(f, lower, upper, CubatureOptions{})
if err == nil || !strings.Contains(err.Error(), "budget") {
t.Fatalf("IntegrateND in 10 dimensions under the default budget: err = %v", err)
}
}
// TestPDEStepCountBound: a dt far below tFinal/1e12 wrapped the step
// count conversion, and the silently larger step ran past the wave
// equation's CFL check.
func TestPDEStepCountBound(t *testing.T) {
u0, err := core.FromFloats([]float64{0, 1, 0, 1, 0}, 5)
if err != nil {
t.Fatal(err)
}
if _, err := IntegrateHeat1D(u0, 1, 0.1, 1, 1e-300, 2, 0, 0); err == nil || !strings.Contains(err.Error(), "1e12") {
t.Fatalf("Heat1D with an unhonourable dt: err = %v", err)
}
u2, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3)
if err != nil {
t.Fatal(err)
}
if _, err := IntegrateHeat2D(u2, 1, 0.1, 0.1, 1, 1e-300, 2, 0, 0, 0, 0); err == nil || !strings.Contains(err.Error(), "1e12") {
t.Fatalf("Heat2D with an unhonourable dt: err = %v", err)
}
}
// TestPDEVerletRejectNonFinite: a NaN or Inf initial state flowed
// through the stencils and published an all-NaN history with no error.
func TestPDEVerletRejectNonFinite(t *testing.T) {
bad, err := core.FromFloats([]float64{1, math.NaN(), 0, 1, 0}, 5)
if err != nil {
t.Fatal(err)
}
if _, err := IntegrateHeat1D(bad, 1, 0.1, 1, 0.1, 2, 0, 0); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("Heat1D on a NaN state: err = %v", err)
}
q, err := core.FromFloats([]float64{1, math.Inf(1)}, 2)
if err != nil {
t.Fatal(err)
}
p, err := core.FromFloats([]float64{0, 0}, 2)
if err != nil {
t.Fatal(err)
}
accel := func(x *core.Array) (*core.Array, error) { return core.Copy(x), nil }
if _, _, err := IntegrateVerlet(accel, 0, 1, q, p, 2); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("Verlet on an Inf state: err = %v", err)
}
}
+363
View File
@@ -0,0 +1,363 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// Adaptive cubature over hyperrectangles: the many-dimensional
// twin of the adaptive Gauss-Legendre quadrature. Each box is measured
// by two product rules (orders 3 and 5 per axis); the difference is
// that box's error estimate, and the globally adaptive loop always
// bisects the worst box along its longest edge, so effort concentrates
// where the integrand actually varies. The 1-D case degenerates to the
// quadrature the package already ships, which doubles as its oracle.
// CubatureOptions tunes IntegrateND. Tolerance bounds the global sum
// of box error estimates (default 1e-10); MaxEvals bounds the function
// evaluations (default two million), an exhausted budget being an
// error naming the achieved estimate, never a silent answer.
type CubatureOptions struct {
Tolerance float64
MaxEvals int
}
// cubBox is one hyperrectangle of the adaptive subdivision: its
// bounds, its measured value and error estimate, and seq, the order
// in which it entered the subdivision. The sequence is the heap's
// tie-break: among equal estimates the earliest inserted box leaves
// first, the same one a scan over the insertion order picks.
type cubBox struct {
lo, hi []float64
val float64
est float64
seq int
}
// cubBoxAbove reports whether a leaves the box heap before b: the
// larger error estimate first, and among equal estimates the earlier
// insertion. Popping that maximum reproduces the selection of a
// linear scan over the insertion order exactly, ties included, for
// every finite estimate. A non-finite estimate can only come out of
// an overflowed measure, a state in which the integral is already
// meaningless: such a box is a total-order special case and stays at
// the bottom of the heap, leaving every finite estimate to run first,
// where the scan would have left it wherever its insertion happened
// to place it. Two non-finite boxes keep insertion order between
// themselves.
func cubBoxAbove(a, b *cubBox) bool {
aOut := math.IsNaN(a.est) || math.IsInf(a.est, 0)
bOut := math.IsNaN(b.est) || math.IsInf(b.est, 0)
if aOut != bOut {
return !aOut
}
if aOut {
return a.seq < b.seq
}
if a.est != b.est {
return a.est > b.est
}
return a.seq < b.seq
}
// cubSiftUp restores the max-heap order after a push at the tail.
func cubSiftUp(h []*cubBox) {
i := len(h) - 1
for i > 0 {
parent := (i - 1) / 2
if !cubBoxAbove(h[i], h[parent]) {
return
}
h[i], h[parent] = h[parent], h[i]
i = parent
}
}
// cubSiftDown restores the max-heap order after the top has been
// replaced from the tail.
func cubSiftDown(h []*cubBox) {
n := len(h)
i := 0
for {
left := 2*i + 1
if left >= n {
return
}
above := left
if right := left + 1; right < n && cubBoxAbove(h[right], h[left]) {
above = right
}
if !cubBoxAbove(h[above], h[i]) {
return
}
h[i], h[above] = h[above], h[i]
i = above
}
}
// cubBoxChunk and cubBoundsChunk size the bisection arenas: one
// allocation per chunk of boxes or bound coordinates instead of one
// per box, so a subdivision that reaches thousands of boxes spends
// tens of allocations, not six per bisection. A chunk never moves once
// handed out, so the heap's pointers stay valid across growth.
const (
cubBoxChunk = 256 // boxes per arena chunk
cubBoundsChunk = 1024 // float64 coordinates per arena chunk
)
// cubBoxArena hands out frozen cubBox values in fixed chunks.
type cubBoxArena struct {
chunks [][]cubBox
}
// alloc returns the next box, zeroed. A box's fields are written once
// by the caller and never after, which is what lets the heap hold the
// pointer for the life of the subdivision.
func (a *cubBoxArena) alloc() *cubBox {
if len(a.chunks) == 0 || len(a.chunks[len(a.chunks)-1]) == cubBoxChunk {
a.chunks = append(a.chunks, make([]cubBox, 0, cubBoxChunk))
}
last := len(a.chunks) - 1
c := append(a.chunks[last], cubBox{})
a.chunks[last] = c
return &c[len(c)-1]
}
// cubBoundsArena hands out box-coordinate slices copied from a parent
// box in fixed chunks. A handed-out slice is written once (the copy,
// then the bisected face) and read-only afterwards.
type cubBoundsArena struct {
chunks [][]float64
used int
}
// copy returns src's values in a fresh arena slice.
func (a *cubBoundsArena) copy(src []float64) []float64 {
n := len(src)
size := max(n, cubBoundsChunk)
if len(a.chunks) == 0 || a.used+n > cap(a.chunks[len(a.chunks)-1]) {
a.chunks = append(a.chunks, make([]float64, 0, size))
a.used = 0
}
last := len(a.chunks) - 1
c := a.chunks[last]
keep := len(c)
c = append(c, src...)
a.chunks[last] = c
a.used += n
return c[keep : keep+n]
}
// IntegrateND returns the integral of f over the hyperrectangle
// [lower, upper] element-wise, by globally adaptive bisection with
// product Gauss-Legendre rules. f receives the evaluation point and
// must not mutate it. A non-finite value, mismatched or empty bounds,
// a reversed edge, or an exhausted evaluation budget is an error.
func IntegrateND(f func(x []float64) float64, lower, upper []float64, opts CubatureOptions) (float64, error) {
const name = "IntegrateND"
if len(lower) == 0 || len(lower) != len(upper) {
return 0, base.Errf("%s: lower and upper must be equal-length non-empty bounds", name)
}
for d := range lower {
if !(upper[d] > lower[d]) {
return 0, base.Errf("%s: edge %d runs from %g to %g", name, d, lower[d], upper[d])
}
}
tol := opts.Tolerance
if tol <= 0 {
tol = 1e-10
}
maxEvals := opts.MaxEvals
if maxEvals <= 0 {
maxEvals = 2_000_000
}
d := len(lower)
n5, w5, err := GaussLegendreNodes(5)
if err != nil {
return 0, base.Errf("%s: %w", name, err)
}
n3, w3, err := GaussLegendreNodes(3)
if err != nil {
return 0, base.Errf("%s: %w", name, err)
}
evals := 0
// One odometer and one evaluation point serve every rule call:
// measure is sequential, so each call overwrites what the last
// read. The bisection budget term is fixed by the dimension.
idx := make([]int, d)
point := make([]float64, d)
// The root box alone costs 5^d + 3^d evaluations before the first
// budget check could fire, and a bisection calls measure twice,
// costing 2·(5^d + 3^d); the loop accounts that true cost below.
// The pre-loop guard is deliberately conservative: it compares
// against the wider bound 8^d + 6^d, saturating, because a MaxInt
// budget must not admit a dimension whose true cost merely fits
// the integer range while needing years to evaluate.
c5, c3, c8, c6 := 1, 1, 1, 1
for range d {
c5 = satMul(c5, 5)
c3 = satMul(c3, 3)
c8 = satMul(c8, 8)
c6 = satMul(c6, 6)
// A saturated product means the true power left the int range:
// it is above every budget, and letting it through would put a
// wrapped count into the later comparisons.
if c8 == math.MaxInt || c6 == math.MaxInt || c8 > maxEvals || c6 > maxEvals {
return 0, base.Errf("%s: dimension %d needs more than the %d-evaluation budget for a single bisection", name, d, maxEvals)
}
}
if c5+c3 > maxEvals {
return 0, base.Errf("%s: dimension %d needs %d evaluations for the root box alone, above the %d budget", name, d, c5+c3, maxEvals)
}
boxEvals := 2*c5 + 2*c3
measure := func(lo, hi []float64) (val, est float64, err error) {
prodRule := func(nodes, weights []float64) (float64, error) {
// One odometer over the per-axis nodes; the axis weights
// multiply along the way, the Jacobian at the end.
jac := 1.0
for a := range d {
jac *= (hi[a] - lo[a]) / 2
}
var sum float64
for {
for a := range d {
point[a] = 0.5*(hi[a]-lo[a])*nodes[idx[a]] + 0.5*(hi[a]+lo[a])
}
w := jac
for a := range d {
w *= weights[idx[a]]
}
v := f(point)
evals++
if math.IsNaN(v) || math.IsInf(v, 0) {
return 0, base.Errf("%s: the integrand is non-finite at %v", name, point)
}
sum += w * v
// Odometer advance.
a := d - 1
for ; a >= 0; a-- {
idx[a]++
if idx[a] < len(nodes) {
break
}
idx[a] = 0
}
if a < 0 {
return sum, nil
}
}
}
fine, err := prodRule(n5, w5)
if err != nil {
return 0, 0, err
}
coarse, err := prodRule(n3, w3)
if err != nil {
return 0, 0, err
}
return fine, math.Abs(fine - coarse), nil
}
rootVal, rootEst, err := measure(lower, upper)
if err != nil {
return 0, err
}
// The boxes awaiting bisection live in a binary max-heap keyed by
// the error estimate with the insertion sequence as the tie-break,
// so, while every estimate stays finite, each pop hands back
// exactly the box a linear scan over the insertion order selects,
// at logarithmic instead of linear cost; an overflowed measure's
// non-finite estimate sorts below every finite one. The sifts work
// index-wise and the backing array grows amortised. The boxes and
// their bound slices come from the chunk arenas above, one
// allocation per chunk instead of per box; the root box aliases the
// caller's bounds, which the solve only reads.
var boxArena cubBoxArena
var boundArena cubBoundsArena
root := boxArena.alloc()
*root = cubBox{lower, upper, rootVal, rootEst, 0}
boxes := []*cubBox{root}
seq := 1
total := rootVal
totalEst := rootEst
// The stopping rule scales the tolerance with the magnitude of the
// integral, the way Integrate combines its bounds: the error
// estimate of an integral of size 1e6 cannot fall below the
// rounding floor of the sum itself, so a purely absolute tolerance
// would burn the whole budget and report an exhausted budget
// instead of the answer. Below unit magnitude the rule is exactly
// the absolute one it always was.
for totalEst > tol*math.Max(1, math.Abs(total)) {
if evals+boxEvals > maxEvals {
return 0, base.Errf("%s: evaluation budget exhausted (%d), estimate %.6g ± %.2g",
name, maxEvals, total, totalEst)
}
// The heap top is the box with the largest error estimate,
// the earliest inserted among equals.
worst := boxes[0]
if worst.est == 0 {
break // every box is already exact by the estimate
}
// Pop it: move the tail box to the top and sift it down.
last := len(boxes) - 1
boxes[0] = boxes[last]
boxes[last] = nil
boxes = boxes[:last]
cubSiftDown(boxes)
// Bisect along the longest edge.
longest := 0
for a := 1; a < d; a++ {
if worst.hi[a]-worst.lo[a] > worst.hi[longest]-worst.lo[longest] {
longest = a
}
}
// The dividing plane keeps every other edge: each child is the
// parent with one face moved to the midpoint, not a corner
// slice (which would collapse the untouched axes). The
// unmodified faces stay the parent's own slices, aliased
// read-only, and the moved face lives in a fresh arena slice:
// box bounds are never written after their one construction
// write, so the aliases hold for the life of the heap.
m := 0.5 * (worst.lo[longest] + worst.hi[longest])
hi1 := boundArena.copy(worst.hi)
hi1[longest] = m
lo2 := boundArena.copy(worst.lo)
lo2[longest] = m
v1, e1, err := measure(worst.lo, hi1)
if err != nil {
return 0, err
}
v2, e2, err := measure(lo2, worst.hi)
if err != nil {
return 0, err
}
b1 := boxArena.alloc()
*b1 = cubBox{worst.lo, hi1, v1, e1, seq}
boxes = append(boxes, b1)
cubSiftUp(boxes)
seq++
b2 := boxArena.alloc()
*b2 = cubBox{lo2, worst.hi, v2, e2, seq}
boxes = append(boxes, b2)
cubSiftUp(boxes)
seq++
total += v1 + v2 - worst.val
totalEst += e1 + e2 - worst.est
}
return total, nil
}
// satMul multiplies with saturation at MaxInt, so a power that
// outgrows the int range reads as "above every budget" instead of
// wrapping into a count the comparisons would read as small.
func satMul(a, b int) int {
if a > math.MaxInt/b {
return math.MaxInt
}
return a * b
}
+112
View File
@@ -0,0 +1,112 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"testing"
)
// cubHeapPop takes the heap's head the way IntegrateND does: the last
// box replaces the root and sifts down.
func cubHeapPop(h []*cubBox) ([]*cubBox, *cubBox) {
top := h[0]
last := len(h) - 1
h[0] = h[last]
h[last] = nil
h = h[:last]
if last > 0 {
cubSiftDown(h)
}
return h, top
}
// TestCubatureBoxHeapOrder pins the order the box heap pops in, the
// substance of the heap replacing the linear scan: the largest finite
// estimate first, the earliest insertion among equals, and a
// non-finite estimate, which only an overflowed measure produces,
// below every finite one.
func TestCubatureBoxHeapOrder(t *testing.T) {
t.Run("larger estimate first", func(t *testing.T) {
h := []*cubBox{{est: 5, seq: 1}, {est: 1, seq: 2}, {est: math.NaN(), seq: 3}, {est: 3, seq: 4}}
for i := range h {
cubSiftUp(h[:i+1])
}
for _, want := range []float64{5, 3, 1} {
var top *cubBox
h, top = cubHeapPop(h)
if top.est != want {
t.Fatalf("popped estimate %v, want %v", top.est, want)
}
}
if h[0].est == h[0].est {
t.Fatalf("a finite estimate %v survived before the non-finite one", h[0].est)
}
})
t.Run("earliest insertion among equals", func(t *testing.T) {
h := []*cubBox{{est: 2, seq: 2}, {est: 2, seq: 0}, {est: 2, seq: 3}, {est: 2, seq: 1}}
for i := range h {
cubSiftUp(h[:i+1])
}
for _, want := range []int{0, 1, 2, 3} {
var top *cubBox
h, top = cubHeapPop(h)
if top.seq != want {
t.Fatalf("popped insertion %d, want %d", top.seq, want)
}
}
})
t.Run("interleaved push and pop", func(t *testing.T) {
h := []*cubBox{{est: 5, seq: 0}}
for _, b := range []*cubBox{{est: 7, seq: 1}, {est: 6, seq: 2}} {
h = append(h, b)
cubSiftUp(h)
}
var top *cubBox
h, top = cubHeapPop(h)
if top.est != 7 {
t.Fatalf("popped estimate %v, want 7", top.est)
}
h = append(h, &cubBox{est: 4, seq: 3})
cubSiftUp(h)
h, top = cubHeapPop(h)
if top.est != 6 {
t.Fatalf("popped estimate %v, want 6", top.est)
}
h, top = cubHeapPop(h)
if top.est != 5 {
t.Fatalf("popped estimate %v, want 5", top.est)
}
})
t.Run("overflowed estimates sort below every finite one", func(t *testing.T) {
// An overflowed measure can carry +Inf, and a corrupted one a
// NaN: both belong at the bottom of the heap, and the insertion
// order holds between them.
h := []*cubBox{
{est: math.Inf(1), seq: 0},
{est: 2, seq: 1},
{est: math.NaN(), seq: 2},
{est: 4, seq: 3},
{est: math.Inf(-1), seq: 4},
{est: 3, seq: 5},
}
for i := range h {
cubSiftUp(h[:i+1])
}
for _, want := range []float64{4, 3, 2} {
var top *cubBox
h, top = cubHeapPop(h)
if top.est != want {
t.Fatalf("popped estimate %v, want %v", top.est, want)
}
}
for _, want := range []int{0, 2, 4} {
var top *cubBox
h, top = cubHeapPop(h)
if top.seq != want {
t.Fatalf("popped insertion %d, want %d", top.seq, want)
}
}
})
}
+114
View File
@@ -0,0 +1,114 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"testing"
)
// TestCubatureGaussian pins the 2-D Gaussian against its exact box
// value π·erf(3)²; the infinite-domain π is not what a box integral
// returns.
func TestCubatureGaussian(t *testing.T) {
got, err := IntegrateND(func(x []float64) float64 {
return math.Exp(-x[0]*x[0] - x[1]*x[1])
}, []float64{-3, -3}, []float64{3, 3}, CubatureOptions{Tolerance: 1e-11})
if err != nil {
t.Fatalf("IntegrateND: %v", err)
}
want := math.Pi * math.Erf(3) * math.Erf(3)
if math.Abs(got-want) > 1e-9 {
t.Fatalf("∫∫e^{-r²} = %.12f, want %.12f", got, want)
}
}
// TestCubaturePolynomials pins exactness on products of polynomials.
func TestCubaturePolynomials(t *testing.T) {
got, err := IntegrateND(func(x []float64) float64 {
return x[0] * x[0] * x[1]
}, []float64{0, 0}, []float64{1, 1}, CubatureOptions{})
if err != nil {
t.Fatalf("IntegrateND: %v", err)
}
if math.Abs(got-1.0/6.0) > 1e-13 {
t.Fatalf("∫x²y = %.14f, want 1/6", got)
}
// 3-D volume of the unit cube shifted.
got3, err := IntegrateND(func(x []float64) float64 { return 1 },
[]float64{1, 2, 3}, []float64{3, 5, 7}, CubatureOptions{})
if err != nil {
t.Fatalf("IntegrateND: %v", err)
}
if math.Abs(got3-24) > 1e-12 {
t.Fatalf("volume = %.12f, want 24", got3)
}
}
// TestCubaturePeaked pins adaptivity: a sharp ridge that uniform
// refinement would crawl on, checked against a dense product Simpson.
func TestCubaturePeaked(t *testing.T) {
f := func(x []float64) float64 {
d2 := (x[0] - 0.4) * (x[0] - 0.4)
d2 += (x[1] - 0.6) * (x[1] - 0.6)
return 1 / (0.003 + d2)
}
got, err := IntegrateND(f, []float64{0, 0}, []float64{1, 1}, CubatureOptions{Tolerance: 1e-9})
if err != nil {
t.Fatalf("IntegrateND: %v", err)
}
// Reference: 800×800 composite midpoint product.
const n = 800
h := 1.0 / n
ref := 0.0
for i := range n {
for j := range n {
ref += h * h * f([]float64{(float64(i) + 0.5) * h, (float64(j) + 0.5) * h})
}
}
if math.Abs(got-ref) > 2e-4*ref {
t.Fatalf("peaked integral = %.8f, reference %.8f", got, ref)
}
}
// TestCubatureMatches1D pins the degenerate dimension against the
// one-dimensional adaptive quadrature.
func TestCubatureMatches1D(t *testing.T) {
f := func(x float64) float64 { return math.Exp(-x) * math.Cos(3*x) }
got, err := IntegrateND(func(x []float64) float64 { return f(x[0]) },
[]float64{0}, []float64{5}, CubatureOptions{Tolerance: 1e-12})
if err != nil {
t.Fatalf("IntegrateND: %v", err)
}
ref, _, err := IntegrateFunction(func(x float64) (float64, error) { return f(x), nil },
0, 5, QuadratureOptions{})
if err != nil {
t.Fatalf("IntegrateFunction: %v", err)
}
if math.Abs(got-ref) > 1e-9 {
t.Fatalf("1-D degenerate = %.12f, quadrature says %.12f", got, ref)
}
}
// TestCubatureErrors pins the input gates.
func TestCubatureErrors(t *testing.T) {
if _, err := IntegrateND(func(x []float64) float64 { return 0 },
[]float64{}, []float64{}, CubatureOptions{}); err == nil {
t.Error("empty bounds accepted")
}
if _, err := IntegrateND(func(x []float64) float64 { return 0 },
[]float64{1, 0}, []float64{0, 1}, CubatureOptions{}); err == nil {
t.Error("reversed edge accepted")
}
if _, err := IntegrateND(func(x []float64) float64 { return math.NaN() },
[]float64{0}, []float64{1}, CubatureOptions{}); err == nil {
t.Error("non-finite integrand accepted")
}
// The budget only bites when refinement is actually needed, so the
// integrand must carry an error estimate a constant cannot.
if _, err := IntegrateND(func(x []float64) float64 { return math.Sin(x[0] * x[1]) },
[]float64{0, 0, 0}, []float64{1, 1, 1}, CubatureOptions{MaxEvals: 1}); err == nil {
t.Error("exhausted budget accepted")
}
}
+109
View File
@@ -0,0 +1,109 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package integrate solves differential equations, integrates functions
// and evolves partial differential equations. It carries five families:
// ordinary differential equations as initial value problems, with event
// detection on the way; two-point boundary value problems by shooting
// and by collocation, and Hamiltonian systems by symplectic schemes;
// adaptive quadrature and cubature; turnkey heat, wave and advection
// solvers in one and two space dimensions; and a piecewise-linear
// finite element Poisson solver on triangular and tetrahedral meshes.
//
// # The state contract
//
// An ordinary differential equation is y' = f(t, y), where f returns the
// derivative of the state at a time and a state. The state is a rank-1
// array: a system of higher rank flattens to its leading-axis vector
// first. The ODE family reads its elements one by one and widens them to
// float64, so an int or float32 state integrates there; the symplectic
// family accepts float64 and float32 positions and momenta only, and a
// complex state is refused everywhere. The returned trajectories are
// float64 arrays, freshly allocated, and the inputs are never written
// to.
//
// Every solver refuses rather than guesses. An exhausted step budget, a
// step size that has collapsed below the resolution of t, an f that
// returns a wrongly shaped state, a non-finite value or a violated CFL
// budget under an explicit stencil is an error naming itself, never a
// silently truncated or silently wrong trajectory.
//
// # Initial value problems
//
// IntegrateODE is the adaptive workhorse (an embedded Dormand-Prince
// 4(5) pair); IntegrateRK4 is the classical fixed-step scheme;
// IntegrateBackwardEuler, IntegrateBDF2, IntegrateBDFVar and
// IntegrateROS4 cover the stiff regime, from the entry-level implicit
// Euler to variable-order BDF and an L-stable Rosenbrock-Wanner method.
// IntegrateODEPath samples the trajectory on an even time grid,
// IntegrateODESteps records every accepted step, and IntegrateODEEvents
// additionally reports where a list of watches crosses zero, filtered by
// direction. IntegrateDAE takes the semi-explicit index-1 mass-matrix
// form M·y' = f(t, y). Backward integration works throughout: a t1 < t0
// integrates in the negative direction.
//
// # Beyond the initial value problem
//
// IntegrateBoundary shoots a two-point boundary value problem, choosing
// the free initial components so the trajectory lands on the prescribed
// end values; SolveBoundaryCollocation solves the same problem by
// three-point Lobatto IIIA collocation on an adaptively refined mesh and
// returns the mesh with the nodal states and slopes. IntegrateVerlet and
// IntegrateYoshida4 integrate a separable Hamiltonian system at a fixed
// step, and IntegrateMidpoint does the same for a general, non-separable
// one through an implicit stage.
//
// # Quadrature and cubature
//
// IntegrateFunction integrates a scalar function over a finite or
// infinite interval and reports an estimate of its own absolute error;
// GaussLegendreNodes hands out the nodes and weights of a fixed rule.
// IntegrateFilon integrates a smooth amplitude against a high-frequency
// sine or cosine carrier, whose cost tracks the amplitude alone rather
// than the carrier a sampled rule must resolve. IntegrateND integrates
// over a hyperrectangle by globally adaptive bisection with product
// Gauss-Legendre rules.
//
// # PDE evolution
//
// IntegrateHeat1D, IntegrateWave1D, IntegrateUpwindAdvection1D,
// IntegrateAdvection1D and IntegrateAdvectionDiffusion1D run on the
// interior grid of a rank-1 initial state, while IntegrateHeat2D and
// IntegrateWave2D run on the rank-2 grid of a rectangle. All of them
// return the trajectory sampled on a time grid, endpoints included.
//
// # Finite elements
//
// GridTriangleMesh2D and BoxTetraMesh3D build structured meshes;
// NewTriangleMesh2D and NewTetraMesh3D accept general conforming ones,
// refusing degenerate elements. SolvePoissonFEM2D and SolvePoissonFEM3D
// assemble and solve -∇·(κ∇u) = f with P1 elements, a conductivity that
// may vary in space, Dirichlet values eliminated by lifting and Neumann
// fluxes integrated on prescribed boundary edges or faces.
//
// # What it deliberately does not do
//
// There is no dense-output object: IntegrateODEPath and
// IntegrateODESteps return the states a caller asked for, and event
// times are narrowed by re-integrating the accepted step rather than
// through a continuous extension. The symplectic family takes a fixed
// step by design, because adaptivity would destroy the property the
// methods exist for. IntegrateDAE is first order and does not project an
// inconsistent start onto the constraint manifold; consistent initial
// values are the caller's contract. The finite element surface is P1 on
// conforming meshes only, and the collocation solver factors a dense
// Newton matrix, so its mesh size is bounded by CollocationOptions. The
// gradient of a trajectory with respect to its parameters is the grad
// package's to compute.
//
// A tour:
//
// end, _ := integrate.IntegrateODE(f, 0, 1, y0, integrate.ODEOptions{})
// hits, end, _ := integrate.IntegrateODEEvents(f, 0, 5, y0, watches, integrate.ODEOptions{})
// area, _ := integrate.IntegrateFunction(g, 0, 1, integrate.QuadratureOptions{})
// history, _ := integrate.IntegrateHeat1D(u0, 1, dx, 0.1, 1e-4, 5, 0, 0)
// u, _ := integrate.SolvePoissonFEM2D(mesh, f, opts) // opts carries κ and the Dirichlet set
//
// The examples in this documentation are executable and checked by the
// test suite.
package integrate
+536
View File
@@ -0,0 +1,536 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The dtype census for integrate: the ODE and PDE drivers' state
// arrays, the FEM mesh tables and the symplectic family, probed with
// Bool, the narrow integers and the Int anchor against a float64
// baseline carrying exactly the widened probe values. The ODE/PDE
// state arrays widen through accessor walks by design, so the narrow
// widths follow Int bit for bit; the symplectic family refuses every
// integer-class state by name, narrow widths included, exactly as it
// refuses Int; the mesh connectivity table keeps its standing Int-only
// gate. Nothing panics or silently misreads.
var igDtypes = []core.Dtype{core.Bool, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32, core.Int}
type igMaker func(vals []float64, shape ...int) *core.Array
func igCast(dt core.Dtype, v float64) float64 {
switch dt {
case core.Bool:
if v != 0 {
return 1
}
return 0
case core.Int8:
return float64(int8(int64(v)))
case core.Uint8:
return float64(uint8(int64(v)))
case core.Int16:
return float64(int16(int64(v)))
case core.Uint16:
return float64(uint16(int64(v)))
case core.Int32:
return float64(int32(int64(v)))
case core.Uint32:
return float64(uint32(int64(v)))
case core.Int:
return float64(int64(v))
default:
return v
}
}
func igMakers(t *testing.T, dt core.Dtype) (probe, base igMaker) {
t.Helper()
castOf := func(vals []float64) []float64 {
out := make([]float64, len(vals))
for i, v := range vals {
out[i] = igCast(dt, v)
}
return out
}
probe = func(vals []float64, shape ...int) *core.Array {
cast := castOf(vals)
var a *core.Array
var err error
switch dt {
case core.Bool:
bs := make([]bool, len(cast))
for i, v := range cast {
bs[i] = v != 0
}
a, err = core.FromBools(bs, shape...)
case core.Int8:
vs := make([]int8, len(cast))
for i, v := range cast {
vs[i] = int8(int64(v))
}
a, err = core.FromInt8s(vs, shape...)
case core.Uint8:
vs := make([]uint8, len(cast))
for i, v := range cast {
vs[i] = uint8(int64(v))
}
a, err = core.FromUint8s(vs, shape...)
case core.Int16:
vs := make([]int16, len(cast))
for i, v := range cast {
vs[i] = int16(int64(v))
}
a, err = core.FromInt16s(vs, shape...)
case core.Uint16:
vs := make([]uint16, len(cast))
for i, v := range cast {
vs[i] = uint16(int64(v))
}
a, err = core.FromUint16s(vs, shape...)
case core.Int32:
vs := make([]int32, len(cast))
for i, v := range cast {
vs[i] = int32(int64(v))
}
a, err = core.FromInt32s(vs, shape...)
case core.Uint32:
vs := make([]uint32, len(cast))
for i, v := range cast {
vs[i] = uint32(int64(v))
}
a, err = core.FromUint32s(vs, shape...)
case core.Int:
vs := make([]int64, len(cast))
for i, v := range cast {
vs[i] = int64(v)
}
a, err = core.FromInts(vs, shape...)
default:
a, err = core.FromFloats(cast, shape...)
}
if err != nil {
t.Fatalf("probe maker (%s): %v", dt, err)
}
return a
}
base = func(vals []float64, shape ...int) *core.Array {
a, err := core.FromFloats(castOf(vals), shape...)
if err != nil {
t.Fatalf("baseline maker: %v", err)
}
return a
}
return probe, base
}
func igElems(t *testing.T, a *core.Array) []float64 {
t.Helper()
out := make([]float64, a.Len())
for i := range out {
out[i] = a.FloatAt(i)
}
return out
}
func igArrays(t *testing.T, label string, dt core.Dtype, probe []*core.Array, perr error, base []*core.Array, berr error) {
t.Helper()
if berr != nil {
if perr == nil {
t.Fatalf("%s(%s): probe succeeded but the float baseline of the same values failed with %v", label, dt, berr)
}
if perr.Error() != berr.Error() {
t.Fatalf("%s(%s): probe error %q differs from the baseline error %q", label, dt, perr, berr)
}
return
}
if perr != nil {
t.Fatalf("%s(%s): %v; the float baseline of the same values succeeded", label, dt, perr)
}
if len(probe) != len(base) {
t.Fatalf("%s(%s): %d outputs against the baseline's %d", label, dt, len(probe), len(base))
}
for k := range probe {
p, b := probe[k], base[k]
if p == nil || b == nil {
t.Fatalf("%s(%s): output %d nil (probe %v, base %v)", label, dt, k, p, b)
}
if p.Dtype() != b.Dtype() {
t.Fatalf("%s(%s): output %d dtype %s, want the baseline dtype %s", label, dt, k, p.Dtype(), b.Dtype())
}
if p.Len() != b.Len() {
t.Fatalf("%s(%s): output %d length %d, want %d", label, dt, k, p.Len(), b.Len())
}
pv, bv := igElems(t, p), igElems(t, b)
for i := range pv {
if pv[i] != bv[i] {
t.Fatalf("%s(%s): output %d element %d = %v, want %v", label, dt, k, i, pv[i], bv[i])
}
}
}
}
func igFloats(t *testing.T, label string, dt core.Dtype, pv []float64, perr error, bv []float64, berr error) {
t.Helper()
if berr != nil {
if perr == nil || perr.Error() != berr.Error() {
t.Fatalf("%s(%s): probe error %v, want the baseline error %v", label, dt, perr, berr)
}
return
}
if perr != nil {
t.Fatalf("%s(%s): %v; the float baseline succeeded", label, dt, perr)
}
if len(pv) != len(bv) {
t.Fatalf("%s(%s): %d values, want %d", label, dt, len(pv), len(bv))
}
for i := range pv {
if pv[i] != bv[i] {
t.Fatalf("%s(%s): value %d = %v, want %v", label, dt, i, pv[i], bv[i])
}
}
}
func igWantErr(t *testing.T, label string, err error, frags ...string) {
t.Helper()
if err == nil {
t.Fatalf("%s: accepted; want a refusal carrying %v", label, frags)
}
for _, f := range frags {
if !strings.Contains(err.Error(), f) {
t.Fatalf("%s: error %q does not contain %q", label, err, f)
}
}
}
// igDecay is the ODE right-hand side the driver rows integrate: the
// callback receives the solver's own float64 state views whatever the
// caller's y0 dtype was.
func igDecay(t float64, y *core.Array) (*core.Array, error) {
out := make([]float64, y.Len())
for i := range out {
out[i] = -y.FloatAt(i)
}
return core.FromFloats(out, len(out))
}
// TestDtypesCensusIntegrate probes every array-taking public entry.
func TestDtypesCensusIntegrate(t *testing.T) {
y0 := []float64{1, 2}
u0 := []float64{1, 2, 3, 4, 5, 6, 7, 8}
u09 := []float64{1, 2, 3, 2, 4, 3, 3, 2, 1}
rows := []struct {
name string
run func(t *testing.T, probe, base igMaker, dt core.Dtype)
}{
{"IntegrateODE", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateODE(igDecay, 0, 1, probe(y0, 2), ODEOptions{})
b, berr := IntegrateODE(igDecay, 0, 1, base(y0, 2), ODEOptions{})
igArrays(t, "IntegrateODE", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateODEPath", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
pt, ps, perr := IntegrateODEPath(igDecay, 0, 1, probe(y0, 2), 4, ODEOptions{})
bt, bs, berr := IntegrateODEPath(igDecay, 0, 1, base(y0, 2), 4, ODEOptions{})
if berr != nil || perr != nil {
igArrays(t, "IntegrateODEPath", dt, nil, perr, nil, berr)
return
}
igFloats(t, "IntegrateODEPath times", dt, pt, nil, bt, nil)
igArrays(t, "IntegrateODEPath states", dt, ps, nil, bs, nil)
}},
{"IntegrateODESteps", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
pt, ps, perr := IntegrateODESteps(igDecay, 0, 1, probe(y0, 2), ODEOptions{})
bt, bs, berr := IntegrateODESteps(igDecay, 0, 1, base(y0, 2), ODEOptions{})
if berr != nil || perr != nil {
igArrays(t, "IntegrateODESteps", dt, nil, perr, nil, berr)
return
}
igFloats(t, "IntegrateODESteps times", dt, pt, nil, bt, nil)
igArrays(t, "IntegrateODESteps states", dt, ps, nil, bs, nil)
}},
{"IntegrateRK4", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateRK4(igDecay, 0, 1, probe(y0, 2), 10)
b, berr := IntegrateRK4(igDecay, 0, 1, base(y0, 2), 10)
igArrays(t, "IntegrateRK4", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateBackwardEuler", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateBackwardEuler(igDecay, 0, 1, probe(y0, 2), 10, ODEOptions{})
b, berr := IntegrateBackwardEuler(igDecay, 0, 1, base(y0, 2), 10, ODEOptions{})
igArrays(t, "IntegrateBackwardEuler", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateBDF2", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateBDF2(igDecay, 0, 1, probe(y0, 2), ODEOptions{})
b, berr := IntegrateBDF2(igDecay, 0, 1, base(y0, 2), ODEOptions{})
igArrays(t, "IntegrateBDF2", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateBDFVar", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateBDFVar(igDecay, 0, 1, probe(y0, 2), BDFVarOptions{})
b, berr := IntegrateBDFVar(igDecay, 0, 1, base(y0, 2), BDFVarOptions{})
igArrays(t, "IntegrateBDFVar", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateROS4", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateROS4(igDecay, 0, 1, probe(y0, 2), ODEOptions{})
b, berr := IntegrateROS4(igDecay, 0, 1, base(y0, 2), ODEOptions{})
igArrays(t, "IntegrateROS4", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateDAE", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
daeF := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{-y.FloatAt(0), y.FloatAt(1)}, 2)
}
p, perr := IntegrateDAE(daeF, probe([]float64{1, 0, 0, 0}, 2, 2), 0, 1, probe([]float64{1, 0}, 2), 5, DAEOptions{})
b, berr := IntegrateDAE(daeF, base([]float64{1, 0, 0, 0}, 2, 2), 0, 1, base([]float64{1, 0}, 2), 5, DAEOptions{})
igArrays(t, "IntegrateDAE", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateODEEvents", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
watch := ODEWatch{Function: func(t float64, y *core.Array) (float64, error) {
return y.FloatAt(0) - 0.5, nil
}}
ph, pf, perr := IntegrateODEEvents(igDecay, 0, 1, probe(y0, 2), []ODEWatch{watch}, ODEOptions{})
bh, bf, berr := IntegrateODEEvents(igDecay, 0, 1, base(y0, 2), []ODEWatch{watch}, ODEOptions{})
if berr != nil || perr != nil {
igArrays(t, "IntegrateODEEvents", dt, nil, perr, nil, berr)
return
}
igArrays(t, "IntegrateODEEvents final", dt, []*core.Array{pf}, nil, []*core.Array{bf}, nil)
if len(ph) != len(bh) {
t.Fatalf("IntegrateODEEvents(%s): %d hits, want %d", dt, len(ph), len(bh))
}
for i := range ph {
if ph[i].Time != bh[i].Time || ph[i].Rising != bh[i].Rising {
t.Fatalf("IntegrateODEEvents(%s): hit %d = (%v, %v), want (%v, %v)",
dt, i, ph[i].Time, ph[i].Rising, bh[i].Time, bh[i].Rising)
}
}
}},
{"IntegrateBoundary", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
osc := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
bc := BoundaryConditions{Start: []int{0}, End: []int{1}, EndValues: []float64{3}}
pt, ps, perr := IntegrateBoundary(osc, 0, 1, probe([]float64{2, 1}, 2), bc, 4, ODEOptions{})
bt, bs, berr := IntegrateBoundary(osc, 0, 1, base([]float64{2, 1}, 2), bc, 4, ODEOptions{})
if berr != nil || perr != nil {
igArrays(t, "IntegrateBoundary", dt, nil, perr, nil, berr)
return
}
igFloats(t, "IntegrateBoundary times", dt, pt, nil, bt, nil)
igArrays(t, "IntegrateBoundary states", dt, ps, nil, bs, nil)
}},
{"SolveBoundaryCollocation", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
osc := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
bc := BoundaryConditions{Start: []int{0}, End: []int{1}, EndValues: []float64{3}}
ps, perr := SolveBoundaryCollocation(osc, 0, 1, probe([]float64{2, 1}, 2), bc, CollocationOptions{})
bs, berr := SolveBoundaryCollocation(osc, 0, 1, base([]float64{2, 1}, 2), bc, CollocationOptions{})
if berr != nil || perr != nil {
igArrays(t, "SolveBoundaryCollocation", dt, nil, perr, nil, berr)
return
}
igFloats(t, "collocation mesh", dt, ps.Mesh, nil, bs.Mesh, nil)
igArrays(t, "collocation values", dt, ps.Values, nil, bs.Values, nil)
}},
{"IntegrateHeat1D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateHeat1D(probe(u0, 8), 1.0, 0.1, 0.01, 0.001, 4, 0, 0)
b, berr := IntegrateHeat1D(base(u0, 8), 1.0, 0.1, 0.01, 0.001, 4, 0, 0)
igArrays(t, "IntegrateHeat1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateWave1D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateWave1D(probe(u0, 8), probe(make([]float64, 8), 8), 0.5, 0.1, 0.01, 0.001, 4)
b, berr := IntegrateWave1D(base(u0, 8), base(make([]float64, 8), 8), 0.5, 0.1, 0.01, 0.001, 4)
igArrays(t, "IntegrateWave1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateHeat2D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateHeat2D(probe(u09, 3, 3), 1.0, 0.5, 0.5, 0.01, 0.001, 3, 0, 0, 0, 0)
b, berr := IntegrateHeat2D(base(u09, 3, 3), 1.0, 0.5, 0.5, 0.01, 0.001, 3, 0, 0, 0, 0)
igArrays(t, "IntegrateHeat2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateWave2D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateWave2D(probe(u09, 3, 3), probe(make([]float64, 9), 3, 3), 0.5, 0.5, 0.5, 0.01, 0.002, 3)
b, berr := IntegrateWave2D(base(u09, 3, 3), base(make([]float64, 9), 3, 3), 0.5, 0.5, 0.5, 0.01, 0.002, 3)
igArrays(t, "IntegrateWave2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateAdvection1D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateAdvection1D(probe(u0, 8), 0.5, 0.1, 0.01, 0.002, 3, 0, 0)
b, berr := IntegrateAdvection1D(base(u0, 8), 0.5, 0.1, 0.01, 0.002, 3, 0, 0)
igArrays(t, "IntegrateAdvection1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateUpwindAdvection1D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateUpwindAdvection1D(probe(u0, 8), 0.5, 0.1, 0.01, 0.002, 3, 0, 0)
b, berr := IntegrateUpwindAdvection1D(base(u0, 8), 0.5, 0.1, 0.01, 0.002, 3, 0, 0)
igArrays(t, "IntegrateUpwindAdvection1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
{"IntegrateAdvectionDiffusion1D", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
p, perr := IntegrateAdvectionDiffusion1D(probe(u0, 8), 0.5, 0.1, 0.1, 0.01, 0.002, 3, 0, 0)
b, berr := IntegrateAdvectionDiffusion1D(base(u0, 8), 0.5, 0.1, 0.1, 0.01, 0.002, 3, 0, 0)
igArrays(t, "IntegrateAdvectionDiffusion1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr)
}},
// The symplectic family refuses integer-class states by name:
// the narrow widths and bool follow Int into the standing
// refusal exactly, the wording unchanged.
{"IntegrateVerlet state gate", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
harmonic := func(q *core.Array) (*core.Array, error) { return core.MulF(q, -1), nil }
_, _, perr := IntegrateVerlet(harmonic, 0, 1, probe([]float64{1}, 1), probe([]float64{0}, 1), 5)
igWantErr(t, "IntegrateVerlet/"+dt.String(), perr,
"IntegrateVerlet", "int states cannot integrate", "float or float32")
}},
{"IntegrateYoshida4 state gate", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
harmonic := func(q *core.Array) (*core.Array, error) { return core.MulF(q, -1), nil }
_, _, perr := IntegrateYoshida4(harmonic, 0, 1, probe([]float64{1}, 1), probe([]float64{0}, 1), 5)
igWantErr(t, "IntegrateYoshida4/"+dt.String(), perr,
"IntegrateYoshida4", "int states cannot integrate", "float or float32")
}},
{"IntegrateMidpoint state gate", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
gradH := func(z *core.Array) (*core.Array, error) {
out := make([]float64, z.Len())
for i := range out {
out[i] = z.FloatAt(i)
}
return core.FromFloats(out, len(out))
}
_, _, perr := IntegrateMidpoint(gradH, 0, 1, probe([]float64{1}, 1), probe([]float64{0}, 1), 5, MidpointOptions{})
igWantErr(t, "IntegrateMidpoint/"+dt.String(), perr,
"IntegrateMidpoint", "int states cannot integrate", "float or float32")
}},
// The callback surface: an acceleration that answers a narrow
// array is read through the accessors exactly as an Int one
// is, with float states the family computes bit-identically.
{"IntegrateVerlet narrow accel output", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
accel := func(q *core.Array) (*core.Array, error) {
return probe([]float64{-q.FloatAt(0)}, 1), nil
}
baseAccel := func(q *core.Array) (*core.Array, error) {
return base([]float64{-q.FloatAt(0)}, 1), nil
}
q0f, _ := core.FromFloats([]float64{1}, 1)
p0f, _ := core.FromFloats([]float64{0}, 1)
pp, pm, perr := IntegrateVerlet(accel, 0, 1, q0f, p0f, 20)
bp, bm, berr := IntegrateVerlet(baseAccel, 0, 1, q0f, p0f, 20)
if berr != nil || perr != nil {
igArrays(t, "Verlet accel output", dt, nil, perr, nil, berr)
return
}
igArrays(t, "Verlet accel positions", dt, pp, nil, bp, nil)
igArrays(t, "Verlet accel momenta", dt, pm, nil, bm, nil)
}},
// FEM mesh tables: vertices widen through accessors (narrow
// follows Int), connectivity keeps the standing Int-only gate.
{"NewTriangleMesh2D vertices widen", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
verts := []float64{0, 0, 1, 0, 0, 1}
tri, terr := core.FromInts([]int64{0, 1, 2}, 1, 3)
if terr != nil {
t.Fatal(terr)
}
p, perr := NewTriangleMesh2D(probe(verts, 3, 2), tri)
b, berr := NewTriangleMesh2D(base(verts, 3, 2), tri)
if berr != nil || perr != nil {
igArrays(t, "NewTriangleMesh2D vertices", dt, nil, perr, nil, berr)
return
}
igFloats(t, "TriangleMesh2D vertices", dt, p.Vertices, nil, b.Vertices, nil)
for i := range p.Triangles {
if p.Triangles[i] != b.Triangles[i] {
t.Fatalf("TriangleMesh2D(%s): triangle %d = %d, want %d", dt, i, p.Triangles[i], b.Triangles[i])
}
}
}},
{"NewTriangleMesh2D connectivity gate", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
if dt == core.Int {
// Int connectivity computes; the gate row covers the
// narrow widths and bool.
return
}
verts, verr := core.FromFloats([]float64{0, 0, 1, 0, 0, 1}, 3, 2)
if verr != nil {
t.Fatal(verr)
}
_, perr := NewTriangleMesh2D(verts, probe([]float64{0, 1, 2}, 1, 3))
igWantErr(t, "NewTriangleMesh2D/"+dt.String(), perr,
"NewTriangleMesh2D", "the triangle table must hold integers", dt.String())
}},
{"NewTetraMesh3D vertices widen", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
verts := []float64{0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1}
tet, terr := core.FromInts([]int64{0, 1, 2, 3}, 1, 4)
if terr != nil {
t.Fatal(terr)
}
p, perr := NewTetraMesh3D(probe(verts, 4, 3), tet)
b, berr := NewTetraMesh3D(base(verts, 4, 3), tet)
if berr != nil || perr != nil {
igArrays(t, "NewTetraMesh3D vertices", dt, nil, perr, nil, berr)
return
}
igFloats(t, "TetraMesh3D vertices", dt, p.Vertices, nil, b.Vertices, nil)
}},
{"NewTetraMesh3D connectivity gate", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
if dt == core.Int {
return
}
verts, verr := core.FromFloats([]float64{0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1}, 4, 3)
if verr != nil {
t.Fatal(verr)
}
_, perr := NewTetraMesh3D(verts, probe([]float64{0, 1, 2, 3}, 1, 4))
igWantErr(t, "NewTetraMesh3D/"+dt.String(), perr,
"NewTetraMesh3D", "the tetrahedron table must hold integers", dt.String())
}},
// The FEM Poisson solvers take no arrays directly: the mesh
// constructor widens the vertex table and gates the connectivity
// table, so a mesh whose vertices carried any probe dtype
// computes identically to the float baseline's.
{"SolvePoissonFEM2D narrow vertices", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
verts := []float64{0, 0, 1, 0, 0, 1}
tri, terr := core.FromInts([]int64{0, 1, 2}, 1, 3)
if terr != nil {
t.Fatal(terr)
}
opts := FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}}
unity := func(x, y float64) float64 { return 1 }
pm, merr := NewTriangleMesh2D(probe(verts, 3, 2), tri)
if merr != nil {
t.Fatalf("NewTriangleMesh2D probe (%s): %v", dt, merr)
}
bm, berr := NewTriangleMesh2D(base(verts, 3, 2), tri)
if berr != nil {
t.Fatalf("NewTriangleMesh2D baseline: %v", berr)
}
p, perr := SolvePoissonFEM2D(pm, unity, opts)
b, serr := SolvePoissonFEM2D(bm, unity, opts)
igArrays(t, "SolvePoissonFEM2D", dt, []*core.Array{p}, perr, []*core.Array{b}, serr)
}},
{"SolvePoissonFEM3D narrow vertices", func(t *testing.T, probe, base igMaker, dt core.Dtype) {
verts := []float64{0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1}
tet, terr := core.FromInts([]int64{0, 1, 2, 3}, 1, 4)
if terr != nil {
t.Fatal(terr)
}
opts := FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}}
unity := func(x, y, z float64) float64 { return 1 }
pm, merr := NewTetraMesh3D(probe(verts, 4, 3), tet)
if merr != nil {
t.Fatalf("NewTetraMesh3D probe (%s): %v", dt, merr)
}
bm, berr := NewTetraMesh3D(base(verts, 4, 3), tet)
if berr != nil {
t.Fatalf("NewTetraMesh3D baseline: %v", berr)
}
p, perr := SolvePoissonFEM3D(pm, unity, opts)
b, serr := SolvePoissonFEM3D(bm, unity, opts)
igArrays(t, "SolvePoissonFEM3D", dt, []*core.Array{p}, perr, []*core.Array{b}, serr)
}},
}
for _, row := range rows {
for _, dt := range igDtypes {
t.Run(row.name+"/"+dt.String(), func(t *testing.T) {
probe, base := igMakers(t, dt)
row.run(t, probe, base, dt)
})
}
}
}
+125
View File
@@ -0,0 +1,125 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression pins for the event and boundary contracts: a watch landing
// exactly on the final boundary, backward event search, the RK4 and
// Verlet refusals of non-finite states, the grid mesh origin screen and
// the cubature budget's true cost.
// TestEventExactlyOnFinalBoundary pins the hit a watch landing
// exactly on zero at the final accepted boundary produces, which the
// sign walk used to swallow.
func TestEventExactlyOnFinalBoundary(t *testing.T) {
f := func(_ float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{1}, 1), nil
}
y0 := mustFloats(t, []float64{0}, 1)
watch := ODEWatch{
Function: func(tt float64, _ *core.Array) (float64, error) { return tt - 1, nil },
Direction: 1,
}
hits, _, err := IntegrateODEEvents(f, 0, 1, y0, []ODEWatch{watch}, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODEEvents: %v", err)
}
if len(hits) != 1 || !hits[0].Rising || math.Abs(hits[0].Time-1) > 1e-12 {
t.Fatalf("hits = %+v, want one rising hit at t = 1", hits)
}
}
// TestEventsBackward pins the event machinery in the backward
// direction: the watch g = t − 0.5 falls through zero at t = 0.5, and
// the backward run must report that hit with the time refined to the
// integrator's accuracy.
func TestEventsBackward(t *testing.T) {
f := func(_ float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{1}, 1), nil
}
y0 := mustFloats(t, []float64{0}, 1)
watch := ODEWatch{
Function: func(tt float64, _ *core.Array) (float64, error) { return tt - 0.5, nil },
Direction: -1,
}
hits, _, err := IntegrateODEEvents(f, 1, 0, y0, []ODEWatch{watch}, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODEEvents backward: %v", err)
}
if len(hits) != 1 || hits[0].Rising {
t.Fatalf("hits = %+v, want one falling hit", hits)
}
if math.Abs(hits[0].Time-0.5) > 1e-9 {
t.Fatalf("hit time = %g, want 0.5", hits[0].Time)
}
}
// TestRK4AndVerletRefuseNonFinite pins the loud refusals on
// the fixed-step integrators, which published NaN states with nil
// errors before.
func TestRK4AndVerletRefuseNonFinite(t *testing.T) {
bad := func(_ float64, _ *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{math.NaN()}, 1), nil
}
y0 := mustFloats(t, []float64{0}, 1)
if _, err := IntegrateRK4(bad, 0, 1, y0, 4); err == nil {
t.Fatal("IntegrateRK4: expected an error for a NaN derivative")
}
accel := func(_ *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{math.Inf(1)}, 1), nil
}
q0 := mustFloats(t, []float64{0}, 1)
p0 := mustFloats(t, []float64{1}, 1)
if _, _, err := IntegrateVerlet(accel, 0, 1, q0, p0, 4); err == nil {
t.Fatal("IntegrateVerlet: expected an error for an Inf acceleration")
}
}
// TestGridMeshRejectsNonFiniteOrigin pins the origin guard.
func TestGridMeshRejectsNonFiniteOrigin(t *testing.T) {
if _, err := GridTriangleMesh2D(math.NaN(), 0, 1, 1, 2, 2); err == nil {
t.Fatal("GridTriangleMesh2D: expected an error for a NaN origin")
}
if _, err := GridTriangleMesh2D(0, math.Inf(1), 1, 1, 2, 2); err == nil {
t.Fatal("GridTriangleMesh2D: expected an error for an Inf origin")
}
// An empty triangle table is refused at construction.
v, _ := core.FromFloats([]float64{0, 0, 1, 0, 0, 1}, 3, 2)
tri, _ := core.FromInts([]int64{}, 0, 3)
if _, err := NewTriangleMesh2D(v, tri); err == nil {
t.Fatal("NewTriangleMesh2D: expected an error for an empty triangle table")
}
}
// TestCubatureBudgetAccountsTrueCost pins the true bisection
// cost 2·(5^d + 3^d): a budget that admits the root box and exactly
// one bisection must complete, and the dimension guard still refuses
// the twenties under any budget.
func TestCubatureBudgetAccountsTrueCost(t *testing.T) {
f := func(x []float64) float64 { return x[0] * x[0] }
lower := []float64{0}
upper := []float64{2}
// Root box: 5 + 3 = 8; one bisection: 2·8 = 16. A budget of 24
// admits the box and one bisection; the old 8^1 + 6^1 = 14
// accounting let the loop overshoot it by two evaluations.
if _, err := IntegrateND(f, lower, upper, CubatureOptions{MaxEvals: 24}); err != nil && !strings.Contains(err.Error(), "converge") {
t.Fatalf("budget 24: err = %v", err)
}
lower25 := make([]float64, 25)
upper25 := make([]float64, 25)
for i := range upper25 {
upper25[i] = 1
}
one := func([]float64) float64 { return 1 }
if _, err := IntegrateND(one, lower25, upper25, CubatureOptions{MaxEvals: math.MaxInt}); err == nil || !strings.Contains(err.Error(), "budget") {
t.Fatalf("d = 25 under a MaxInt budget: err = %v", err)
}
}
+236
View File
@@ -0,0 +1,236 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression pins: the event detector's blind first step, watch
// values that signed themselves across zero, PDE parameters that were
// only half guarded, the step recorder's rounded endpoint, and a
// cubature budget the dimension powers could switch off.
// TestEventFirstAcceptedStep: the detector seeded its comparison from
// the END of the first accepted step, so a crossing inside that step
// went unnoticed; the seed is now the watch value at the step's start
// state, and the crossing is refined like any other.
func TestEventFirstAcceptedStep(t *testing.T) {
zero := func(now float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{0}, 1), nil
}
cross := func(now float64, y *core.Array) (float64, error) {
return now - 0.0005, nil
}
hits, _, err := IntegrateODEEvents(zero, 0, 1, mustFloats(t, []float64{1}, 1),
[]ODEWatch{{Function: cross}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents: %v", err)
}
if len(hits) != 1 {
t.Fatalf("hits = %d, want exactly the crossing near 0.0005", len(hits))
}
if math.Abs(hits[0].Time-0.0005) > 1e-9 {
t.Fatalf("hit at %.14g, want 0.0005", hits[0].Time)
}
}
// TestEventWatchNonFiniteRefused: a NaN from a watch compared false in
// the sign test and manufactured a crossing (or swallowed one); a
// non-finite watch value is now an error naming the value.
func TestEventWatchNonFiniteRefused(t *testing.T) {
zero := func(now float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{0}, 1), nil
}
calls := 0
nanFirst := func(now float64, y *core.Array) (float64, error) {
calls++
if calls == 1 {
return math.NaN(), nil
}
return -1, nil
}
_, _, err := IntegrateODEEvents(zero, 0, 1, mustFloats(t, []float64{1}, 1),
[]ODEWatch{{Function: nanFirst}}, ODEOptions{})
if err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("a NaN watch value: err = %v", err)
}
infLater := func(now float64, y *core.Array) (float64, error) {
if now > 0.2 {
return math.Inf(-1), nil
}
return -1, nil
}
_, _, err = IntegrateODEEvents(zero, 0, 1, mustFloats(t, []float64{1}, 1),
[]ODEWatch{{Function: infLater}}, ODEOptions{})
if err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("an infinite watch value: err = %v", err)
}
}
// TestGridMeshFiniteExtents: a NaN or Inf extent passed the old
// positivity test (a NaN compares false against <= 0) and laid out a
// mesh of non-finite vertices.
func TestGridMeshFiniteExtents(t *testing.T) {
for name, extents := range map[string][2]float64{
"NaN width": {math.NaN(), 1},
"NaN height": {1, math.NaN()},
"Inf width": {math.Inf(1), 1},
"Inf height": {1, math.Inf(-1)},
} {
mesh, err := GridTriangleMesh2D(0, 0, extents[0], extents[1], 2, 2)
if err == nil || !strings.Contains(err.Error(), "finite") {
t.Fatalf("%s: err = %v, mesh = %v", name, err, mesh != nil)
}
}
// A valid grid still builds.
if _, err := GridTriangleMesh2D(0, 0, 1, 1, 2, 2); err != nil {
t.Fatalf("a valid grid: %v", err)
}
}
// TestFEMConstantKappaNonFinite: a +Inf constant conductivity slipped
// through the positivity test and died mid-factorisation.
func TestFEMConstantKappaNonFinite(t *testing.T) {
mesh, err := GridTriangleMesh2D(0, 0, 1, 1, 4, 4)
if err != nil {
t.Fatalf("GridTriangleMesh2D: %v", err)
}
opts := FEMPoissonOptions{Kappa: math.Inf(1), DirichletNodes: []int{0}, DirichletValues: []float64{0}}
if _, err := SolvePoissonFEM2D(mesh, nil, opts); err == nil || !strings.Contains(err.Error(), "positive") {
t.Fatalf("a +Inf constant conductivity: err = %v", err)
}
// With the conductivity field set, a non-finite placeholder for
// the constant is refused all the same.
opts.KappaFunc = func(x, y float64) float64 { return 1 }
if _, err := SolvePoissonFEM2D(mesh, nil, opts); err == nil || !strings.Contains(err.Error(), "positive") {
t.Fatalf("a +Inf placeholder conductivity beside KappaFunc: err = %v", err)
}
}
// TestPDEParameterNonFiniteRefusals walks the solvers' numeric
// parameters: each one used to slip a NaN or Inf past a comparison
// that reads false against NaN and publish an all-NaN history.
func TestPDEParameterNonFiniteRefusals(t *testing.T) {
u1, v1 := mustFloats(t, []float64{0, 1, 0, 1, 0}, 5), mustFloats(t, []float64{0, 0, 0, 0, 0}, 5)
u2 := mustFloats(t, []float64{0, 1, 0, 0, 1, 0, 0, 1, 0}, 3, 3)
cases := []struct {
name string
run func() (*core.Array, error)
}{
{"Heat1D kappa +Inf", func() (*core.Array, error) {
return IntegrateHeat1D(u1, math.Inf(1), 0.1, 0.1, 0.01, 2, 0, 0)
}},
{"Heat1D kappa NaN", func() (*core.Array, error) {
return IntegrateHeat1D(u1, math.NaN(), 0.1, 0.1, 0.01, 2, 0, 0)
}},
{"Heat1D NaN bound", func() (*core.Array, error) {
return IntegrateHeat1D(u1, 1, 0.1, 0.1, 0.01, 2, math.NaN(), 0)
}},
{"Heat1D Inf bound", func() (*core.Array, error) {
return IntegrateHeat1D(u1, 1, 0.1, 0.1, 0.01, 2, 0, math.Inf(1))
}},
{"Wave1D c NaN", func() (*core.Array, error) {
return IntegrateWave1D(u1, v1, math.NaN(), 0.1, 0.1, 0.01, 2)
}},
{"Wave1D c Inf", func() (*core.Array, error) {
return IntegrateWave1D(u1, v1, math.Inf(1), 0.1, 0.1, 0.01, 2)
}},
{"Wave1D v0 NaN", func() (*core.Array, error) {
return IntegrateWave1D(u1, mustFloats(t, []float64{0, math.NaN(), 0, 0, 0}, 5), 1, 0.1, 0.1, 0.01, 2)
}},
{"Heat2D kappa +Inf", func() (*core.Array, error) {
return IntegrateHeat2D(u2, math.Inf(1), 0.1, 0.1, 0.1, 0.01, 2, 0, 0, 0, 0)
}},
{"Heat2D NaN boundary", func() (*core.Array, error) {
return IntegrateHeat2D(u2, 1, 0.1, 0.1, 0.1, 0.01, 2, 0, 0, math.NaN(), 0)
}},
{"Heat2D Inf boundary", func() (*core.Array, error) {
return IntegrateHeat2D(u2, 1, 0.1, 0.1, 0.1, 0.01, 2, 0, 0, 0, math.Inf(1))
}},
{"Wave2D v0 NaN", func() (*core.Array, error) {
return IntegrateWave2D(u2, mustFloats(t, []float64{0, 0, 0, 0, math.NaN(), 0, 0, 0, 0}, 3, 3), 1, 0.1, 0.1, 0.1, 0.01, 2)
}},
}
for _, c := range cases {
if _, err := c.run(); err == nil || !strings.Contains(err.Error(), "finite") && !strings.Contains(err.Error(), "positive") {
t.Fatalf("%s: err = %v, want a finite/positive refusal", c.name, err)
}
}
}
// TestODEStepsEndpointExact: the recorder's last time was the run's
// accumulated t+h, a few ulps off t1; it is now t1 exactly.
func TestODEStepsEndpointExact(t *testing.T) {
decayF := func(now float64, y *core.Array) (*core.Array, error) {
return core.MulF(y, -1), nil
}
times, states, err := IntegrateODESteps(decayF, 0, 0.3, mustFloats(t, []float64{1}, 1), ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODESteps: %v", err)
}
if last := times[len(times)-1]; last != 0.3 {
t.Fatalf("last recorded time = %.17g, want 0.3 exactly", last)
}
// The pinned endpoint still closes on the analytic curve.
if last := states[len(states)-1].FloatAt(0); math.Abs(last-math.Exp(-0.3)) > 1e-6 {
t.Fatalf("y(0.3) = %.14g, want %.14g", last, math.Exp(-0.3))
}
// Across magnitudes the run's own boundary misses t1 by whole
// ulps (t0 = 1e16 has an ulp of 2 and the span is 2): the endpoint
// is pinned regardless, and the recorded state there is the
// answer IntegrateODE itself returns for t1.
flat := func(now float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(0)}, 1)
}
const (
big = 1e16
span = 2.0
)
times, states, err = IntegrateODESteps(flat, big, big+span, mustFloats(t, []float64{3}, 1), ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODESteps across magnitudes: %v", err)
}
if last := times[len(times)-1]; last != big+span {
t.Fatalf("last recorded time = %.17g, want %.17g exactly", last, big+span)
}
if got := states[len(states)-1].FloatAt(0); got != 3 {
t.Fatalf("state at the endpoint = %g, want the constant 3", got)
}
}
// TestCubatureDimensionPowerSaturates: past the twenties the straight
// int powers wrapped, the bisection cost went negative and every
// budget check with it; the powers now saturate and a saturated power
// reads as above the budget.
func TestCubatureDimensionPowerSaturates(t *testing.T) {
lower := make([]float64, 25)
upper := make([]float64, 25)
for i := range upper {
upper[i] = 1
}
f := func(x []float64) float64 { return 1 }
// A small budget is refused on the single-bisection cost alone.
if _, err := IntegrateND(f, lower, upper, CubatureOptions{MaxEvals: 1024}); err == nil || !strings.Contains(err.Error(), "budget") {
t.Fatalf("d = 25 under a 1024-evaluation budget: err = %v", err)
}
// A budget of MaxInt used to walk the powers straight into the
// wrap and then evaluate the 5^25-point root box: the saturated
// computation refuses it before the first evaluation.
if _, err := IntegrateND(f, lower, upper, CubatureOptions{MaxEvals: math.MaxInt}); err == nil || !strings.Contains(err.Error(), "budget") {
t.Fatalf("d = 25 under a MaxInt budget: err = %v", err)
}
// A sane dimension and budget still integrate.
got, err := IntegrateND(f, lower[:3], upper[:3], CubatureOptions{})
if err != nil {
t.Fatalf("d = 3 under the default budget: %v", err)
}
if math.Abs(got-1) > 1e-10 {
t.Fatalf("integral of 1 over the unit cube = %.17g, want 1", got)
}
}
+225
View File
@@ -0,0 +1,225 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate_test
// Runnable examples for the package: the flagship workflows, each
// with a fixed output that `go test` checks, so the printed
// documentation cannot drift from the code.
import (
"fmt"
"log"
"math"
tensor "sourcedock.dev/petrbalvin/tensor"
"sourcedock.dev/petrbalvin/tensor/integrate"
)
// The stiff scalar problem y' = −1000·(y − cos t) − sin t with
// y(0) = 1, whose exact solution is y = cos t. The variable-order,
// variable-step BDF scheme takes the long steps the solution's
// smoothness allows where a fixed small step would be forced by the
// fast transient, and BDFVarStats reports what it did.
func ExampleIntegrateBDFVar() {
f := func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{-1000*(y.FloatAt(0)-math.Cos(t)) - math.Sin(t)}, 1)
}
y0, _ := tensor.FromFloats([]float64{1}, 1)
var stats integrate.BDFVarStats
y, err := integrate.IntegrateBDFVar(f, 0, 1, y0, integrate.BDFVarOptions{Stats: &stats})
if err != nil {
log.Fatal(err)
}
fmt.Printf("y(1) = %.6f, exact %.6f\n", y.FloatAt(0), math.Cos(1))
fmt.Printf("accepted %d steps, rejected %d, highest order %d\n", stats.Steps, stats.Rejected, stats.MaxOrder)
// Output:
// y(1) = 0.540302, exact 0.540302
// accepted 21 steps, rejected 1, highest order 5
}
// Event detection along a trajectory: the oscillator y″ = −y started
// at y = (1, 0) passes the level y = 0.5 falling at t = π/3 and rising
// at t = 5π/3. Each watch carries its own direction filter, and
// IntegrateODEEvents returns the crossings sorted by time alongside
// the final state.
func ExampleIntegrateODEEvents() {
f := func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
y0, _ := tensor.FromFloats([]float64{1, 0}, 2)
// The same level function twice, with opposite direction filters;
// Direction 0 would record both crossings on one watch.
level := func(t float64, y *tensor.Array) (float64, error) { return y.FloatAt(0) - 0.5, nil }
watches := []integrate.ODEWatch{
{Function: level, Direction: -1},
{Function: level, Direction: +1},
}
hits, final, err := integrate.IntegrateODEEvents(f, 0, 7, y0, watches, integrate.ODEOptions{})
if err != nil {
log.Fatal(err)
}
for _, h := range hits {
direction := "falling"
if h.Rising {
direction = "rising"
}
fmt.Printf("watch %d fired at t = %.4f (%s), y = %.4f\n", h.Watch, h.Time, direction, h.State.FloatAt(0))
}
fmt.Printf("y(7) = %.4f\n", final.FloatAt(0))
// Output:
// watch 0 fired at t = 1.0472 (falling), y = 0.5000
// watch 1 fired at t = 5.2360 (rising), y = 0.5000
// y(7) = 0.7539
}
// A symplectic integrator on a separable Hamiltonian: the harmonic
// oscillator q″ = −q with unit mass, q(0) = 1 and p(0) = 0, whose
// energy ½(p² + q²) stays in a bounded band instead of drifting.
// The step stays fixed by design; only the number of steps is chosen.
func ExampleIntegrateVerlet() {
accel := func(q *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{-q.FloatAt(0)}, 1)
}
q0, _ := tensor.FromFloats([]float64{1}, 1)
p0, _ := tensor.FromFloats([]float64{0}, 1)
const steps = 4000
positions, momenta, err := integrate.IntegrateVerlet(accel, 0, 2*math.Pi, q0, p0, steps)
if err != nil {
log.Fatal(err)
}
worst := 0.0
for s := range steps + 1 {
q, p := positions[s].FloatAt(0), momenta[s].FloatAt(0)
drift := math.Abs(0.5*(p*p+q*q) - 0.5)
worst = math.Max(worst, drift)
}
fmt.Printf("q(2π) = %.6f, p(2π) = %.2e\n", positions[steps].FloatAt(0), momenta[steps].FloatAt(0))
fmt.Printf("worst energy deviation over the period: %.2e\n", worst)
// Output:
// q(2π) = 1.000000, p(2π) = -6.46e-07
// worst energy deviation over the period: 3.08e-07
}
// Quadrature and cubature: a Gauss-Legendre rule read from the
// package's cache, an adaptive integral over an infinite range, and a
// two-dimensional integral by globally adaptive bisection.
func Example_quadratureAndCubature() {
nodes, weights, err := integrate.GaussLegendreNodes(3)
if err != nil {
log.Fatal(err)
}
for i := range nodes {
fmt.Printf("node %.6f, weight %.6f\n", nodes[i], weights[i])
}
value, errEst, err := integrate.IntegrateFunction(func(x float64) (float64, error) {
return math.Exp(-x * x), nil
}, 0, math.Inf(1), integrate.QuadratureOptions{})
if err != nil {
log.Fatal(err)
}
fmt.Printf("the Gaussian tail integrates to %.6f (error estimate %.1e)\n", value, errEst)
area, err := integrate.IntegrateND(func(x []float64) float64 {
return x[0] * x[1]
}, []float64{0, 0}, []float64{1, 1}, integrate.CubatureOptions{})
if err != nil {
log.Fatal(err)
}
fmt.Printf("x·y over the unit square integrates to %.6f\n", area)
// Output:
// node -0.774597, weight 0.555556
// node 0.000000, weight 0.888889
// node 0.774597, weight 0.555556
// the Gaussian tail integrates to 0.886227 (error estimate 9.8e-12)
// x·y over the unit square integrates to 0.250000
}
// Heat evolution in one dimension: u_t = u_xx on [0, 1] from
// u = sin(πx), Dirichlet ends held at zero. The sampled history is a
// (samples, n) array of interior states, and the centre decays as the
// exact e^(−π²t)·sin(π/2) predicts.
func ExampleIntegrateHeat1D() {
const n, samples = 399, 5
const dx, tFinal = 1.0 / 400, 0.1
u0 := make([]float64, n)
for i := range u0 {
u0[i] = math.Sin(math.Pi * float64(i+1) * dx)
}
state, err := tensor.FromFloats(u0, n)
if err != nil {
log.Fatal(err)
}
history, err := integrate.IntegrateHeat1D(state, 1, dx, tFinal, 1e-4, samples, 0, 0)
if err != nil {
log.Fatal(err)
}
centre := history.FloatAt((samples-1)*n + n/2)
exact := math.Exp(-math.Pi * math.Pi * tFinal)
fmt.Printf("history shape %v\n", history.Shape())
fmt.Printf("u(1/2, 0.1) = %.6f, exact %.6f\n", centre, exact)
// Output:
// history shape [5 399]
// u(1/2, 0.1) = 0.372710, exact 0.372708
}
// The finite element Poisson solve: −∇·(κ∇u) = f on the unit square
// with κ = 1 and Dirichlet data on the boundary ring. The manufactured
// solution u = sin(πx)·sin(πy) makes f = 2π²·sin(πx)·sin(πy), and the
// P1 solution reproduces it to the mesh's accuracy at the centre.
func ExampleSolvePoissonFEM2D() {
const cells = 16
mesh, err := integrate.GridTriangleMesh2D(0, 0, 1, 1, cells, cells)
if err != nil {
log.Fatal(err)
}
var nodes []int
var values []float64
for v := range mesh.Vertices2() {
x, y := mesh.Vertices[2*v], mesh.Vertices[2*v+1]
onEdge := x == 0 || x == 1 || y == 0 || y == 1
if onEdge {
nodes = append(nodes, v)
values = append(values, math.Sin(math.Pi*x)*math.Sin(math.Pi*y))
}
}
u, err := integrate.SolvePoissonFEM2D(mesh, func(x, y float64) float64 {
return 2 * math.Pi * math.Pi * math.Sin(math.Pi*x) * math.Sin(math.Pi*y)
}, integrate.FEMPoissonOptions{Kappa: 1, DirichletNodes: nodes, DirichletValues: values})
if err != nil {
log.Fatal(err)
}
centre := (cells/2)*(cells+1) + cells/2
fmt.Printf("mesh of %d vertices, %d triangles, %d boundary edges\n",
mesh.Vertices2(), mesh.Triangles3(), len(mesh.BoundaryEdges())/2)
fmt.Printf("u(1/2, 1/2) = %.4f on this mesh, exact 1.0000\n", u.FloatAt(centre))
// Output:
// mesh of 289 vertices, 512 triangles, 64 boundary edges
// u(1/2, 1/2) = 0.9946 on this mesh, exact 1.0000
}
// The two-point boundary value problem: y″ = −y with y(0) = 0 and
// y(π/2) = 1, solved by shooting on the free initial slope. The slope
// comes out as 1 and the sampled trajectory traces y = sin t.
func ExampleIntegrateBoundary() {
f := func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
y0, _ := tensor.FromFloats([]float64{0, 0.5}, 2)
bc := integrate.BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}}
times, states, err := integrate.IntegrateBoundary(f, 0, math.Pi/2, y0, bc, 3,
integrate.ODEOptions{RelTol: 1e-10, AbsTol: 1e-13})
if err != nil {
log.Fatal(err)
}
fmt.Printf("shooting slope y'(0) = %.6f\n", states[0].FloatAt(1))
for i := range times {
fmt.Printf("y(%.4f) = %.6f, exact %.6f\n", times[i], states[i].FloatAt(0), math.Sin(times[i]))
}
// Output:
// shooting slope y'(0) = 1.000000
// y(0.0000) = 0.000000, exact 0.000000
// y(0.7854) = 0.707107, exact 0.707107
// y(1.5708) = 1.000000, exact 1.000000
}
+376
View File
@@ -0,0 +1,376 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"slices"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
linalg "sourcedock.dev/petrbalvin/tensor/linalg"
)
// The finite element surface for second-order problems on general
// two-dimensional domains: piecewise-linear (P1) elements on a
// conforming triangular mesh, the stiffness matrix assembled straight
// into the sparse triple format, Dirichlet values eliminated by
// lifting, Neumann boundaries free of charge, and the reduced system
// handed to the sparse Cholesky factorisation the direct-solvers
// surface provides.
// TriangleMesh2D carries a conforming triangular mesh: vertex
// coordinates as x,y pairs and triangles as triples of vertex
// indices. The orientation of a triangle does not matter; a triangle
// with zero area does and is refused at construction.
type TriangleMesh2D struct {
// Vertices holds x,y for every vertex: two entries per vertex.
Vertices []float64
// Triangles holds three vertex indices per triangle.
Triangles []int64
}
// NewTriangleMesh2D builds a mesh from a vertex table with two
// columns and a triangle table with three columns of vertex indices.
// Indices must lie in range and a degenerate triangle (three
// collinear vertices) is an error: its stiffness contribution is
// undefined.
func NewTriangleMesh2D(vertices *core.Array, triangles *core.Array) (*TriangleMesh2D, error) {
const name = "NewTriangleMesh2D"
if vertices.Dtype() == core.Complex || triangles.Dtype() == core.Complex {
return nil, base.Errf("%s: complex mesh data is not supported", name)
}
if vertices.NDim() != 2 || vertices.Shape()[1] != 2 {
return nil, base.Errf("%s: the vertex table must be rank 2 with two columns, got shape %s", name, base.ShapeText(vertices.Shape()))
}
if triangles.Dtype() != core.Int {
return nil, base.Errf("%s: the triangle table must hold integers, got %s", name, triangles.Dtype())
}
if triangles.NDim() != 2 || triangles.Shape()[1] != 3 {
return nil, base.Errf("%s: the triangle table must be rank 2 with three columns, got shape %s", name, base.ShapeText(triangles.Shape()))
}
n := vertices.Shape()[0]
m := triangles.Shape()[0]
if n < 3 {
return nil, base.Errf("%s: a mesh needs at least three vertices, got %d", name, n)
}
if m == 0 {
// An empty triangle table would surface deep in the sparse
// factorisation on the zero rows of the free nodes, far from
// the mesh that caused it.
return nil, base.Errf("%s: the triangle table must not be empty", name)
}
mesh := &TriangleMesh2D{Vertices: make([]float64, 2*n), Triangles: make([]int64, 3*m)}
for i := range 2 * n {
v := vertices.FloatAt(i)
if math.IsNaN(v) || math.IsInf(v, 0) {
return nil, base.Errf("%s: vertex coordinate %d is not finite", name, i)
}
mesh.Vertices[i] = v
}
for p := range 3 * m {
idx := triangles.RawInts()[p]
if idx < 0 || idx >= int64(n) {
return nil, base.Errf("%s: triangle vertex index %d out of range for %d vertices", name, idx, n)
}
mesh.Triangles[p] = idx
}
// A triangle with zero area carries no stiffness: refuse it here
// where the caller can name the triangle, not mid-assembly.
for t := range m {
a, b, c := mesh.Triangles[3*t], mesh.Triangles[3*t+1], mesh.Triangles[3*t+2]
ax, ay := mesh.Vertices[2*a], mesh.Vertices[2*a+1]
bx, by := mesh.Vertices[2*b], mesh.Vertices[2*b+1]
cx, cy := mesh.Vertices[2*c], mesh.Vertices[2*c+1]
if area := math.Abs((bx-ax)*(cy-ay)-(cx-ax)*(by-ay)) / 2; area == 0 {
return nil, base.Errf("%s: triangle %d is degenerate (zero area)", name, t)
}
}
return mesh, nil
}
// Vertices2 returns the vertex count.
func (m *TriangleMesh2D) Vertices2() int { return len(m.Vertices) / 2 }
// Triangles3 returns the triangle count.
func (m *TriangleMesh2D) Triangles3() int { return len(m.Triangles) / 3 }
// BoundaryEdges returns the mesh's boundary edges as flat pairs of
// vertex indices: an edge belongs to the boundary when exactly one
// triangle carries it. The pairs are sorted, so the result is a pure
// function of the mesh.
func (m *TriangleMesh2D) BoundaryEdges() []int {
count := make(map[[2]int]int, len(m.Triangles))
key := func(a, b int) [2]int {
if a < b {
return [2]int{a, b}
}
return [2]int{b, a}
}
for t := 0; t < m.Triangles3(); t++ {
a, b, c := int(m.Triangles[3*t]), int(m.Triangles[3*t+1]), int(m.Triangles[3*t+2])
count[key(a, b)]++
count[key(b, c)]++
count[key(c, a)]++
}
edges := make([]int, 0, 8)
for e, n := range count {
if n == 1 {
edges = append(edges, e[0], e[1])
}
}
slices.Sort(edges)
return edges
}
// GridTriangleMesh2D builds the structured triangulation of the
// axis-aligned rectangle [x0, x0+width] × [y0, y0+height] with m by n
// cells, two triangles per cell. m and n must both be positive.
func GridTriangleMesh2D(x0, y0, width, height float64, m, n int) (*TriangleMesh2D, error) {
const name = "GridTriangleMesh2D"
if m <= 0 || n <= 0 {
return nil, base.Errf("%s: the cell counts must be positive, got %d by %d", name, m, n)
}
// The same guard NewTriangleMesh2D applies to its vertex table: a
// non-finite extent or origin would lay out vertices at NaN or Inf
// and only surface mid-factorisation, far from the cause.
if !(width > 0) || !(height > 0) || math.IsInf(width, 0) || math.IsInf(height, 0) ||
math.IsNaN(x0) || math.IsInf(x0, 0) || math.IsNaN(y0) || math.IsInf(y0, 0) {
return nil, base.Errf("%s: the extents must be finite and positive and the origin finite, got origin (%g, %g), extents %g by %g",
name, x0, y0, width, height)
}
vertices := make([]float64, 2*(m+1)*(n+1))
for j := range n + 1 {
for i := range m + 1 {
vertices[2*(j*(m+1)+i)] = x0 + width*float64(i)/float64(m)
vertices[2*(j*(m+1)+i)+1] = y0 + height*float64(j)/float64(n)
}
}
at := func(i, j int) int64 { return int64(j*(m+1) + i) }
triangles := make([]int64, 0, 6*m*n)
for j := range n {
for i := range m {
triangles = append(triangles,
at(i, j), at(i+1, j), at(i+1, j+1),
at(i, j), at(i+1, j+1), at(i, j+1))
}
}
return &TriangleMesh2D{Vertices: vertices, Triangles: triangles}, nil
}
// FEMPoissonOptions carries the data SolvePoissonFEM2D needs beside
// the mesh and the source: the conductivity, the prescribed boundary
// values, and the optional flux boundary.
type FEMPoissonOptions struct {
// Kappa is the constant conductivity when KappaFunc is nil. It
// must be positive.
Kappa float64
// KappaFunc, when set, gives the conductivity at a point. It is
// evaluated at the triangle centroids and must be positive there
// for every triangle; a non-positive value names the triangle.
KappaFunc func(x, y float64) float64
// DirichletNodes lists the vertices with prescribed values and
// DirichletValues the values in the same order. The nodes leave
// the system with their rows and columns; at least one is
// required, because a purely Neumann problem has no unique
// solution.
DirichletNodes []int
DirichletValues []float64
// NeumannEdges lists boundary edges as flat pairs of vertex
// indices and NeumannFlux gives the flux κ∂u/∂n along each edge's
// outward normal: each edge receives half of length·flux at its
// midpoint into both endpoints. A nil flux means zero.
NeumannEdges []int
NeumannFlux func(x, y float64) float64
// Ordering selects the fill-reducing permutation for the sparse
// Cholesky factorisation. The zero value is the natural order;
// meshes usually want SparseOrderingReverseCuthillMcKee.
Ordering linalg.SparseOrdering
}
// SolvePoissonFEM2D solves −∇·(κ∇u) = f on the mesh with
// piecewise-linear elements: the stiffness matrix is assembled per
// triangle (the conductivity evaluated at the centroids when it
// varies), the load is lumped at the vertices from f at the
// centroids, Neumann fluxes are integrated along their edges, and
// Dirichlet values are eliminated by lifting. f may be nil for the
// homogeneous equation.
func SolvePoissonFEM2D(mesh *TriangleMesh2D, f func(x, y float64) float64, opts FEMPoissonOptions) (*core.Array, error) {
const name = "SolvePoissonFEM2D"
if mesh == nil {
return nil, base.Errf("%s: the mesh must not be nil", name)
}
n := mesh.Vertices2()
// With KappaFunc nil the constant conductivity is the value used,
// so it must be positive and finite; with the field set the
// constant is a placeholder, but a non-finite one is still refused
// rather than silently ignored.
if opts.KappaFunc == nil {
if !(opts.Kappa > 0) || math.IsInf(opts.Kappa, 0) {
return nil, base.Errf("%s: the conductivity must be positive, got %g", name, opts.Kappa)
}
} else if math.IsNaN(opts.Kappa) || math.IsInf(opts.Kappa, 0) {
return nil, base.Errf("%s: the conductivity must be positive, got %g", name, opts.Kappa)
}
if len(opts.DirichletNodes) != len(opts.DirichletValues) {
return nil, base.Errf("%s: %d Dirichlet nodes but %d values", name, len(opts.DirichletNodes), len(opts.DirichletValues))
}
if len(opts.DirichletNodes) == 0 {
return nil, base.Errf("%s: a purely Neumann problem has no unique solution; prescribe at least one Dirichlet value", name)
}
// The Dirichlet nodes as a dense marker with their prescribed
// values: the lifting and the unit rows below each visit every
// assembled entry, and a marker answers those visits in constant
// time where a set of nodes answered with a hash. A node listed
// twice keeps its last value and appears once, as it did in the
// set; the appended order does not reach the assembled system,
// whose coordinate entries the sparse conversion sorts and merges
// by coordinate.
dirichletMark := make([]bool, n)
dirichletVal := make([]float64, n)
dirichletNodes := make([]int, 0, len(opts.DirichletNodes))
for p, d := range opts.DirichletNodes {
if d < 0 || d >= n {
return nil, base.Errf("%s: Dirichlet node %d out of range for %d vertices", name, d, n)
}
v := opts.DirichletValues[p]
if math.IsNaN(v) || math.IsInf(v, 0) {
return nil, base.Errf("%s: Dirichlet value at node %d is not finite", name, d)
}
if !dirichletMark[d] {
dirichletNodes = append(dirichletNodes, d)
}
dirichletMark[d] = true
dirichletVal[d] = v
}
if len(opts.NeumannEdges)%2 != 0 {
return nil, base.Errf("%s: %d Neumann edge indices, want pairs", name, len(opts.NeumannEdges))
}
for p := 0; p < len(opts.NeumannEdges); p += 2 {
a, b := opts.NeumannEdges[p], opts.NeumannEdges[p+1]
if a < 0 || a >= n || b < 0 || b >= n || a == b {
return nil, base.Errf("%s: Neumann edge [%d,%d] is not a valid vertex pair", name, a, b)
}
}
// Assembly: nine entries per triangle, symmetric by construction;
// the load is lumped one third of the triangle area to each of
// its vertices, with the conductivity evaluated at the centroid
// when it varies.
entries := make([]float64, 0, 9*mesh.Triangles3())
rows := make([]int, 0, 9*mesh.Triangles3())
cols := make([]int, 0, 9*mesh.Triangles3())
load := make([]float64, n)
for t := 0; t < mesh.Triangles3(); t++ {
a, b, c := int(mesh.Triangles[3*t]), int(mesh.Triangles[3*t+1]), int(mesh.Triangles[3*t+2])
ax, ay := mesh.Vertices[2*a], mesh.Vertices[2*a+1]
bx, by := mesh.Vertices[2*b], mesh.Vertices[2*b+1]
cx, cy := mesh.Vertices[2*c], mesh.Vertices[2*c+1]
area := math.Abs((bx-ax)*(cy-ay)-(cx-ax)*(by-ay)) / 2
kappa := opts.Kappa
if opts.KappaFunc != nil {
kappa = opts.KappaFunc((ax+bx+cx)/3, (ay+by+cy)/3)
if !(kappa > 0) || math.IsNaN(kappa) || math.IsInf(kappa, 0) {
return nil, base.Errf("%s: the conductivity at triangle %d is %g, want positive", name, t, kappa)
}
}
// The gradient basis: b are the y differences, c the x
// differences, and K = κ/(4A)·(b⊗b + c⊗c).
bb := [3]float64{by - cy, cy - ay, ay - by}
cc := [3]float64{cx - bx, ax - cx, bx - ax}
nodes := [3]int{a, b, c}
for i := range 3 {
for j := range 3 {
v := kappa * (bb[i]*bb[j] + cc[i]*cc[j]) / (4 * area)
rows = append(rows, nodes[i])
cols = append(cols, nodes[j])
entries = append(entries, v)
}
}
if f != nil {
fv := f((ax+bx+cx)/3, (ay+by+cy)/3)
// A non-finite source value would flow into the load and the
// solve would publish an all-NaN solution with a nil error,
// the breach every other integrator here refuses up front.
if math.IsNaN(fv) || math.IsInf(fv, 0) {
return nil, base.Errf("%s: the source returned the non-finite value %g at triangle %d", name, fv, t)
}
contribution := area / 3 * fv
load[a] += contribution
load[b] += contribution
load[c] += contribution
}
}
// Neumann fluxes: half of length·flux into each endpoint of every
// listed edge, the flux evaluated at the edge midpoint.
if len(opts.NeumannEdges) > 0 && opts.NeumannFlux != nil {
for p := 0; p < len(opts.NeumannEdges); p += 2 {
a, b := opts.NeumannEdges[p], opts.NeumannEdges[p+1]
ax, ay := mesh.Vertices[2*a], mesh.Vertices[2*a+1]
bx, by := mesh.Vertices[2*b], mesh.Vertices[2*b+1]
length := math.Hypot(bx-ax, by-ay)
fv := opts.NeumannFlux((ax+bx)/2, (ay+by)/2)
// A non-finite flux lands in the load like a non-finite
// source, so the same refusal answers it.
if math.IsNaN(fv) || math.IsInf(fv, 0) {
return nil, base.Errf("%s: the Neumann flux returned the non-finite value %g on edge [%d, %d]", name, fv, a, b)
}
flux := length / 2 * fv
load[a] += flux
load[b] += flux
}
}
// Dirichlet lifting: the known boundary values move to the right
// hand side, then their rows and columns leave the system as
// unit rows.
for p, i := range rows {
if j := cols[p]; dirichletMark[j] {
load[i] -= entries[p] * dirichletVal[j]
}
}
keptRows := make([]int64, 0, len(rows))
keptCols := make([]int64, 0, len(rows))
keptVals := make([]float64, 0, len(rows))
for p := range rows {
i, j := rows[p], cols[p]
if dirichletMark[i] || dirichletMark[j] {
continue
}
keptRows = append(keptRows, int64(i))
keptCols = append(keptCols, int64(j))
keptVals = append(keptVals, entries[p])
}
for _, d := range dirichletNodes {
keptRows = append(keptRows, int64(d))
keptCols = append(keptCols, int64(d))
keptVals = append(keptVals, 1)
load[d] = dirichletVal[d]
}
indices, err := core.FromInts(pairInts(keptRows, keptCols), len(keptVals), 2)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
coo, err := core.NewSparseCOO(indices, fromSlice(keptVals, len(keptVals)), []int{n, n})
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
order := opts.Ordering
factor, err := linalg.NewSparseCholesky(coo, order)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
rhs := core.New(core.Float, []int{n}...)
copy(rhs.RawFloats(), load)
return factor.Solve(rhs)
}
// pairInts interleaves row and column indices into the index table
// the sparse coordinate format expects.
func pairInts(rows, cols []int64) []int64 {
out := make([]int64, 2*len(rows))
for p := range rows {
out[2*p] = rows[p]
out[2*p+1] = cols[p]
}
return out
}
+477
View File
@@ -0,0 +1,477 @@
// 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"
linalg "sourcedock.dev/petrbalvin/tensor/linalg"
)
// gridMesh builds the structured triangulation of the unit square
// with m cells per side, two triangles per cell, and returns the mesh
// plus the list of boundary vertices in the order (bottom row, top
// row, left column, right column), duplicates removed.
func gridMesh(t *testing.T, m int) (*TriangleMesh2D, []int) {
t.Helper()
mesh, err := GridTriangleMesh2D(0, 0, 1, 1, m, m)
if err != nil {
t.Fatalf("GridTriangleMesh2D: %v", err)
}
boundary := make([]int, 0, 4*m)
for i := range m + 1 {
boundary = append(boundary, i, m*(m+1)+i)
}
for j := 1; j < m; j++ {
boundary = append(boundary, j*(m+1), j*(m+1)+m)
}
return mesh, boundary
}
func TestSolvePoissonFEM2DConvergence(t *testing.T) {
// The manufactured solution u = sin(πx)·sin(πy) on the unit
// square drives f = 2π²·u; with the boundary lifted the P1 error
// must halve twice when the mesh is refined, the O(h²) the
// piecewise-linear theory promises.
solution := func(x, y float64) float64 { return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) }
source := func(x, y float64) float64 { return 2 * math.Pi * math.Pi * solution(x, y) }
previous := 0.0
for _, m := range []int{8, 16, 32} {
mesh, boundary := gridMesh(t, m)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[2*node], mesh.Vertices[2*node+1])
}
u, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM2D(m=%d): %v", m, err)
}
worst := 0.0
for i := range mesh.Vertices2() {
if d := math.Abs(u.FloatAt(i) - solution(mesh.Vertices[2*i], mesh.Vertices[2*i+1])); d > worst {
worst = d
}
}
t.Logf("m=%2d: max nodal error %.3g", m, worst)
if previous > 0 && previous/worst < 2.5 {
t.Fatalf("m=%d: refinement ratio %.2f, want the O(h²) rate (previous %.3g, now %.3g)",
m, previous/worst, previous, worst)
}
if m == 32 && worst > 2e-3 {
t.Fatalf("m=32: error %.3g too large for the asymptotic range", worst)
}
previous = worst
}
}
// TestSolvePoissonFEM2DLinearExactness is the patch test the P1
// elements must pass without compromise: a linear field lies in the
// approximation space, so with f = 0 and the boundary lifted the
// interior solution must equal the field to machine precision.
func TestSolvePoissonFEM2DLinearExactness(t *testing.T) {
mesh, boundary := gridMesh(t, 12)
field := func(x, y float64) float64 { return 1 + 2*x - 3*y }
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = field(mesh.Vertices[2*node], mesh.Vertices[2*node+1])
}
u, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM2D: %v", err)
}
worst := 0.0
for i := range mesh.Vertices2() {
if d := math.Abs(u.FloatAt(i) - field(mesh.Vertices[2*i], mesh.Vertices[2*i+1])); d > worst {
worst = d
}
}
if worst > 1e-12 {
t.Fatalf("linear patch test error %.3g, want machine precision", worst)
}
}
// TestSolvePoissonFEM2DNeumannNatural pins the natural boundary: a
// constant field with f = 0 satisfies the homogeneous Neumann
// condition everywhere, so pinning the constant at a single vertex
// must reproduce it across the whole mesh.
func TestSolvePoissonFEM2DNeumannNatural(t *testing.T) {
mesh, _ := gridMesh(t, 10)
u, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{5}})
if err != nil {
t.Fatalf("SolvePoissonFEM2D: %v", err)
}
for i := range mesh.Vertices2() {
if math.Abs(u.FloatAt(i)-5) > 1e-10 {
t.Fatalf("node %d: solution %.12g, want the constant 5", i, u.FloatAt(i))
}
}
}
// TestSolvePoissonFEM2DOrderings runs the manufactured-solution solve
// under every ordering the factor offers: the ordering changes the
// fill, never the answer.
func TestSolvePoissonFEM2DOrderings(t *testing.T) {
solution := func(x, y float64) float64 { return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) }
mesh, boundary := gridMesh(t, 10)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[2*node], mesh.Vertices[2*node+1])
}
source := func(x, y float64) float64 { return 2 * math.Pi * math.Pi * solution(x, y) }
reference, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM2D(natural): %v", err)
}
for _, ordering := range []linalg.SparseOrdering{
linalg.SparseOrderingReverseCuthillMcKee,
linalg.SparseOrderingMinimumDegree,
} {
u, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values, Ordering: ordering})
if err != nil {
t.Fatalf("SolvePoissonFEM2D(%d): %v", ordering, err)
}
for i := range mesh.Vertices2() {
if math.Abs(u.FloatAt(i)-reference.FloatAt(i)) > 1e-9 {
t.Fatalf("ordering %d: node %d differs from the natural run", ordering, i)
}
}
}
}
func TestSolvePoissonFEM2DRefusals(t *testing.T) {
mesh, boundary := gridMesh(t, 5)
// Degenerate triangle: three collinear vertices.
if _, err := NewTriangleMesh2D(
floatsToArrayFEM(t, []float64{0, 0, 1, 0, 2, 0}, 3, 2),
intsToArrayFEM(t, []int64{0, 1, 2}, 1, 3)); err == nil || !stringsContains(err, "degenerate") {
t.Fatalf("a degenerate triangle: %v", err)
}
// Triangle index out of range.
if _, err := NewTriangleMesh2D(
floatsToArrayFEM(t, []float64{0, 0, 1, 0, 0, 1}, 3, 2),
intsToArrayFEM(t, []int64{0, 1, 3}, 1, 3)); err == nil || !stringsContains(err, "out of range") {
t.Fatalf("out of range index: %v", err)
}
// A float triangle table: the triangles must be integer indices.
if _, err := NewTriangleMesh2D(
floatsToArrayFEM(t, []float64{0, 0, 1, 0, 0, 1}, 3, 2),
floatsToArrayFEM(t, []float64{0, 1, 2}, 1, 3)); err == nil || !stringsContains(err, "integers") {
t.Fatalf("a float triangle table: %v", err)
}
// Dirichlet node out of range and a length mismatch.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{99}, DirichletValues: []float64{1}}); err == nil || !stringsContains(err, "out of range") {
t.Fatalf("an out of range Dirichlet node: %v", err)
}
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0, 1}, DirichletValues: []float64{1}}); err == nil {
t.Fatal("a Dirichlet length mismatch was accepted")
}
// Non-positive conductivity.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 0, DirichletNodes: boundary, DirichletValues: make([]float64, len(boundary))}); err == nil || !stringsContains(err, "positive") {
t.Fatalf("zero conductivity: %v", err)
}
// Non-finite Dirichlet value.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{math.NaN()}}); err == nil || !stringsContains(err, "finite") {
t.Fatalf("a NaN Dirichlet value: %v", err)
}
// A NaN vertex coordinate in the mesh table.
if _, err := NewTriangleMesh2D(
floatsToArrayFEM(t, []float64{math.NaN(), 0, 1, 0, 0, 1}, 3, 2),
intsToArrayFEM(t, []int64{0, 1, 2}, 1, 3)); err == nil || !stringsContains(err, "not finite") {
t.Fatalf("a NaN vertex coordinate: %v", err)
}
// An odd number of Neumann edge indices: no complete pairs.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}, NeumannEdges: []int{0, 1, 2}}); err == nil || !stringsContains(err, "pairs") {
t.Fatalf("an odd Neumann edge count: %v", err)
}
// A degenerate Neumann edge a == b.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}, NeumannEdges: []int{3, 3}}); err == nil || !stringsContains(err, "valid vertex pair") {
t.Fatalf("a degenerate Neumann edge: %v", err)
}
// A Neumann edge index out of range.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}, NeumannEdges: []int{0, 999}}); err == nil || !stringsContains(err, "valid vertex pair") {
t.Fatalf("an out of range Neumann edge: %v", err)
}
// A KappaFunc returning a non-positive conductivity names the
// triangle instead of assembling a singular stiffness matrix.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{
KappaFunc: func(float64, float64) float64 { return -1 },
DirichletNodes: []int{0},
DirichletValues: []float64{0},
}); err == nil || !stringsContains(err, "positive") {
t.Fatalf("a non-positive KappaFunc value: %v", err)
}
// A KappaFunc returning an infinite conductivity names the triangle
// the way the constant field's gate names itself, instead of
// surfacing as a factorisation failure far from the cause.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{
KappaFunc: func(float64, float64) float64 { return math.Inf(1) },
DirichletNodes: []int{0},
DirichletValues: []float64{0},
}); err == nil || !stringsContains(err, "positive") {
t.Fatalf("an infinite KappaFunc value: %v", err)
}
// A non-finite source value refuses the solve: it used to land in
// the load and publish an all-NaN solution with a nil error.
if _, err := SolvePoissonFEM2D(mesh, func(x, y float64) float64 { return math.NaN() },
FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: make([]float64, len(boundary))}); err == nil || !stringsContains(err, "non-finite") {
t.Fatalf("a NaN source value: %v", err)
}
// A non-finite Neumann flux refuses the solve the same way.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{
Kappa: 1,
DirichletNodes: boundary,
DirichletValues: make([]float64, len(boundary)),
NeumannEdges: []int{0, mesh.Vertices2() - 1},
NeumannFlux: func(x, y float64) float64 { return math.Inf(1) },
}); err == nil || !stringsContains(err, "non-finite") {
t.Fatalf("an infinite Neumann flux: %v", err)
}
// An ordering that does not exist.
if _, err := SolvePoissonFEM2D(mesh, nil, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: make([]float64, len(boundary)), Ordering: linalg.SparseOrdering(7)}); err == nil {
t.Fatal("an unknown ordering was accepted")
}
}
// TestSolvePoissonFEM2DIsDeterministic solves the same problem twice
// and requires bit-identical nodal values, the contract every Tensor
// entry point carries.
func TestSolvePoissonFEM2DIsDeterministic(t *testing.T) {
solution := func(x, y float64) float64 { return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) }
mesh, boundary := gridMesh(t, 10)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[2*node], mesh.Vertices[2*node+1])
}
source := func(x, y float64) float64 { return 2 * math.Pi * math.Pi * solution(x, y) }
u1, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("first solve: %v", err)
}
u2, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("second solve: %v", err)
}
for i := range mesh.Vertices2() {
if u1.FloatAt(i) != u2.FloatAt(i) {
t.Fatalf("node %d differs: %.17g vs %.17g", i, u1.FloatAt(i), u2.FloatAt(i))
}
}
}
func stringsContains(err error, fragment string) bool {
return err != nil && len(err.Error()) >= len(fragment) && indexOf(err.Error(), fragment) >= 0
}
func indexOf(s, fragment string) int {
for i := 0; i+len(fragment) <= len(s); i++ {
if s[i:i+len(fragment)] == fragment {
return i
}
}
return -1
}
func floatsToArrayFEM(t *testing.T, vals []float64, shape ...int) *core.Array {
t.Helper()
a, err := core.FromFloats(vals, shape...)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
return a
}
func intsToArrayFEM(t *testing.T, vals []int64, shape ...int) *core.Array {
t.Helper()
a, err := core.FromInts(vals, shape...)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
return a
}
// TestSolvePoissonFEM2DNeumannFlux pins the boundary-edge integrals:
// u = (x²+y²)/2 has −Δu = −2 and the flux κ∂u/∂n = 1 along the right
// and top edges' outward normals (0 along the bottom and left), so
// prescribing those fluxes with a single pinned vertex must
// reproduce the quadratic field. The midpoint edge rule is
// first-order consistent, so the error must halve with the mesh.
func TestSolvePoissonFEM2DNeumannFlux(t *testing.T) {
field := func(x, y float64) float64 { return (x*x + y*y) / 2 }
previous := 0.0
for _, m := range []int{10, 20} {
mesh, err := GridTriangleMesh2D(0, 0, 1, 1, m, m)
if err != nil {
t.Fatalf("GridTriangleMesh2D: %v", err)
}
// Boundary edges: pairs of neighbouring boundary vertices.
var edges []int
at := func(i, j int) int { return j*(m+1) + i }
for j := range m {
edges = append(edges, at(j, 0), at(j+1, 0)) // bottom: flux 0
edges = append(edges, at(j, m), at(j+1, m)) // top: flux 1
edges = append(edges, at(m, j), at(m, j+1)) // right: flux 1
edges = append(edges, at(0, j), at(0, j+1)) // left: flux 0
}
flux := func(x, y float64) float64 {
if x == 1 || y == 1 {
return 1
}
return 0
}
u, err := SolvePoissonFEM2D(mesh, func(float64, float64) float64 { return -2 },
FEMPoissonOptions{
Kappa: 1,
DirichletNodes: []int{at(0, 0)},
DirichletValues: []float64{0},
NeumannEdges: edges,
NeumannFlux: flux,
})
if err != nil {
t.Fatalf("m=%d: %v", m, err)
}
worst := 0.0
for i := range mesh.Vertices2() {
x := mesh.Vertices[2*i]
y := mesh.Vertices[2*i+1]
if d := math.Abs(u.FloatAt(i) - field(x, y)); d > worst {
worst = d
}
}
t.Logf("m=%2d: max nodal error %.3g", m, worst)
if previous > 0 && previous/worst < 1.4 {
t.Fatalf("m=%d: refinement ratio %.2f, want the first-order flux rate", m, previous/worst)
}
if m == 20 && worst > 5e-3 {
t.Fatalf("m=20: error %.3g too large", worst)
}
previous = worst
}
}
// TestSolvePoissonFEM2DVariableKappa runs the manufactured solution
// with a spatially varying conductivity evaluated at the element
// centroids: f must carry the analytic divergence terms, and the
// P1 convergence rate must survive the varying coefficient.
func TestSolvePoissonFEM2DVariableKappa(t *testing.T) {
sin, cos := math.Pi, math.Pi
u := func(x, y float64) float64 { return math.Sin(sin*x) * math.Sin(sin*y) }
kappaF := func(x, y float64) float64 { return 1 + x*y }
ux := func(x, y float64) float64 { return cos * math.Cos(cos*x) * math.Sin(cos*y) }
uy := func(x, y float64) float64 { return cos * math.Sin(cos*x) * math.Cos(cos*y) }
lap := func(x, y float64) float64 { return -2 * math.Pi * math.Pi * u(x, y) }
source := func(x, y float64) float64 {
k := kappaF(x, y)
return -(y*ux(x, y) + x*uy(x, y) + k*lap(x, y))
}
previous := 0.0
for _, m := range []int{8, 16, 32} {
mesh, boundary := gridMesh(t, m)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = u(mesh.Vertices[2*node], mesh.Vertices[2*node+1])
}
uk, err := SolvePoissonFEM2D(mesh, source,
FEMPoissonOptions{KappaFunc: kappaF, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM2D(m=%d): %v", m, err)
}
worst := 0.0
for i := range mesh.Vertices2() {
if d := math.Abs(uk.FloatAt(i) - u(mesh.Vertices[2*i], mesh.Vertices[2*i+1])); d > worst {
worst = d
}
}
t.Logf("m=%2d: max nodal error %.3g", m, worst)
if previous > 0 && previous/worst < 2.5 {
t.Fatalf("m=%d: refinement ratio %.2f, want the O(h²) rate", m, previous/worst)
}
previous = worst
}
}
// TestTriangleMesh2DBoundaryEdges pins the boundary-edge detection:
// the m by n grid carries exactly 2(m+n) boundary edges, every one of
// them with both endpoints on the boundary vertex ring.
func TestTriangleMesh2DBoundaryEdges(t *testing.T) {
mesh, err := GridTriangleMesh2D(0, 0, 1, 1, 5, 3)
if err != nil {
t.Fatalf("GridTriangleMesh2D: %v", err)
}
edges := mesh.BoundaryEdges()
if len(edges) != 2*2*(5+3) {
t.Fatalf("boundary edge count %d, want %d", len(edges), 2*(5+3))
}
onBoundary := func(v int) bool {
i := v % 6
j := v / 6
return i == 0 || i == 5 || j == 0 || j == 3
}
for p := 0; p < len(edges); p += 2 {
if !onBoundary(edges[p]) || !onBoundary(edges[p+1]) {
t.Fatalf("edge [%d,%d] is not on the boundary", edges[p], edges[p+1])
}
}
// The generator's vertex positions are exact.
mesh2, err := GridTriangleMesh2D(-1, 2, 2, 4, 2, 2)
if err != nil {
t.Fatalf("GridTriangleMesh2D: %v", err)
}
if mesh2.Vertices[0] != -1 || mesh2.Vertices[1] != 2 {
t.Fatalf("vertex 0 = [%g %g], want [-1 2]", mesh2.Vertices[0], mesh2.Vertices[1])
}
if mesh2.Vertices[2*(2*3+2)] != 1 || mesh2.Vertices[2*(2*3+2)+1] != 6 {
t.Fatalf("vertex (2,2) = [%g %g], want [1 6]",
mesh2.Vertices[2*(2*3+2)], mesh2.Vertices[2*(2*3+2)+1])
}
if _, err := GridTriangleMesh2D(0, 0, 1, 1, 0, 3); err == nil {
t.Fatal("a zero cell count was accepted")
}
if _, err := GridTriangleMesh2D(0, 0, -1, 1, 2, 2); err == nil {
t.Fatal("a negative extent was accepted")
}
}
// TestSolvePoissonFEM2DDuplicateDirichletNode pins the documented rule
// for a node listed more than once: the last value is the prescribed
// one and the node enters the assembled system exactly once, so the
// repeated listing answers what the single listing with that value
// answers. Recording it twice appends a second unit row at the same
// coordinate, which the sparse conversion merges by summing, so the
// node's diagonal doubles and the solve halves its prescribed value.
func TestSolvePoissonFEM2DDuplicateDirichletNode(t *testing.T) {
solution := func(x, y float64) float64 { return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) }
source := func(x, y float64) float64 { return 2 * math.Pi * math.Pi * solution(x, y) }
mesh, boundary := gridMesh(t, 4)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[2*node], mesh.Vertices[2*node+1])
}
// The list names the second boundary node again at the end, with a
// different value: the last one wins and the node stays single.
const extra = 0.5
nodes := append(append([]int(nil), boundary...), boundary[1])
dupValues := append(append([]float64(nil), values...), values[1]+extra)
u, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: nodes, DirichletValues: dupValues})
if err != nil {
t.Fatalf("SolvePoissonFEM2D with a repeated node: %v", err)
}
single := append([]float64(nil), values...)
single[1] += extra
want, err := SolvePoissonFEM2D(mesh, source, FEMPoissonOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: single})
if err != nil {
t.Fatalf("SolvePoissonFEM2D with the node once: %v", err)
}
if got := u.FloatAt(boundary[1]); math.Abs(got-(values[1]+extra)) > 1e-12 {
t.Fatalf("the repeated node answered %g, want the last prescribed value %g", got, values[1]+extra)
}
worst := 0.0
for i := range mesh.Vertices2() {
worst = math.Max(worst, math.Abs(u.FloatAt(i)-want.FloatAt(i)))
}
if worst > 1e-12 {
t.Fatalf("the repeated listing differs from the single listing by %g, want the node recorded once", worst)
}
}
+21
View File
@@ -0,0 +1,21 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import "testing"
// TestSolvePoissonFEM2DNilMeshRefused pins the nil-mesh refusal: the
// three-dimensional solve reports a nil mesh as an error, so the
// two-dimensional one answers the same way instead of dereferencing
// it and panicking.
func TestSolvePoissonFEM2DNilMeshRefused(t *testing.T) {
_, err := SolvePoissonFEM2D(nil, nil, FEMPoissonOptions{
Kappa: 1,
DirichletNodes: []int{0},
DirichletValues: []float64{0},
})
if err == nil {
t.Fatal("SolvePoissonFEM2D: expected an error for a nil mesh")
}
}
+552
View File
@@ -0,0 +1,552 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"fmt"
"math"
"slices"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
linalg "sourcedock.dev/petrbalvin/tensor/linalg"
)
// The finite element groundwork for second-order problems in three
// dimensions, the volumetric sibling of the triangular surface in
// fem2d.go: piecewise-linear (P1) elements on a conforming
// tetrahedral mesh, the stiffness matrix assembled per tetrahedron
// from the gradient-of-basis formula over the element's edge vectors,
// the load integrated per element with a collapsed Gauss rule,
// Dirichlet values eliminated by lifting, Neumann fluxes integrated
// on prescribed boundary faces, and the reduced system handed to the
// same sparse Cholesky factorisation the two-dimensional path uses.
// TetraMesh3D carries a conforming tetrahedral mesh: vertex
// coordinates as x,y,z triples and tetrahedra as quadruples of vertex
// indices in positive orientation, meaning the signed volume
// (b−a)·((c−a)×(d−a)) of every stored tetrahedron is positive. A
// tetrahedron with zero volume or negative orientation does matter
// and is refused at construction.
type TetraMesh3D struct {
// Vertices holds x,y,z for every vertex: three entries per vertex.
Vertices []float64
// Tetrahedra holds four vertex indices per tetrahedron.
Tetrahedra []int64
}
// NewTetraMesh3D builds a mesh from a vertex table with three columns
// and a tetrahedron table with four columns of vertex indices.
// Indices must lie in range, every coordinate must be finite, and a
// degenerate (zero-volume) or inverted (negative-orientation)
// tetrahedron is an error naming the element and its vertices: its
// stiffness contribution is undefined.
func NewTetraMesh3D(vertices *core.Array, tetrahedra *core.Array) (*TetraMesh3D, error) {
const name = "NewTetraMesh3D"
if vertices.Dtype() == core.Complex || tetrahedra.Dtype() == core.Complex {
return nil, base.Errf("%s: complex mesh data is not supported", name)
}
if vertices.NDim() != 2 || vertices.Shape()[1] != 3 {
return nil, base.Errf("%s: the vertex table must be rank 2 with three columns, got shape %s",
name, base.ShapeText(vertices.Shape()))
}
if tetrahedra.Dtype() != core.Int {
return nil, base.Errf("%s: the tetrahedron table must hold integers, got %s", name, tetrahedra.Dtype())
}
if tetrahedra.NDim() != 2 || tetrahedra.Shape()[1] != 4 {
return nil, base.Errf("%s: the tetrahedron table must be rank 2 with four columns, got shape %s",
name, base.ShapeText(tetrahedra.Shape()))
}
n := vertices.Shape()[0]
m := tetrahedra.Shape()[0]
if n < 4 {
return nil, base.Errf("%s: a mesh needs at least four vertices, got %d", name, n)
}
if m == 0 {
// An empty tetrahedron table would surface deep in the sparse
// factorisation on the zero rows of the free nodes, far from
// the mesh that caused it.
return nil, base.Errf("%s: the tetrahedron table must not be empty", name)
}
mesh := &TetraMesh3D{Vertices: make([]float64, 3*n), Tetrahedra: make([]int64, 4*m)}
for i := range 3 * n {
v := vertices.FloatAt(i)
if math.IsNaN(v) || math.IsInf(v, 0) {
return nil, base.Errf("%s: vertex coordinate %d is not finite", name, i)
}
mesh.Vertices[i] = v
}
for q := range 4 * m {
idx := tetrahedra.RawInts()[q]
if idx < 0 || idx >= int64(n) {
return nil, base.Errf("%s: tetrahedron vertex index %d out of range for %d vertices", name, idx, n)
}
mesh.Tetrahedra[q] = idx
}
// Orientation and volume are checked where the caller can name the
// tetrahedron and its vertices, not mid-assembly. Both messages
// carry the coordinates, so a mis-ordered table can be fixed
// without reopening a mesh debugger.
for t := range m {
a, b, c, d := int(mesh.Tetrahedra[4*t]), int(mesh.Tetrahedra[4*t+1]), int(mesh.Tetrahedra[4*t+2]), int(mesh.Tetrahedra[4*t+3])
ax, ay, az := mesh.Vertices[3*a], mesh.Vertices[3*a+1], mesh.Vertices[3*a+2]
bx, by, bz := mesh.Vertices[3*b], mesh.Vertices[3*b+1], mesh.Vertices[3*b+2]
cx, cy, cz := mesh.Vertices[3*c], mesh.Vertices[3*c+1], mesh.Vertices[3*c+2]
dx, dy, dz := mesh.Vertices[3*d], mesh.Vertices[3*d+1], mesh.Vertices[3*d+2]
signed6 := signedTetraVolume(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz)
at := func(v int) string {
return fmt.Sprintf("(%g, %g, %g)", mesh.Vertices[3*v], mesh.Vertices[3*v+1], mesh.Vertices[3*v+2])
}
verts := fmt.Sprintf("vertices %d %s, %d %s, %d %s, %d %s", a, at(a), b, at(b), c, at(c), d, at(d))
if signed6 == 0 {
return nil, base.Errf("%s: tetrahedron %d is degenerate (zero volume), %s", name, t, verts)
}
if signed6 < 0 {
return nil, base.Errf("%s: tetrahedron %d is inverted (signed volume %g), %s", name, t, signed6/6, verts)
}
}
return mesh, nil
}
// Vertices3 returns the vertex count.
func (m *TetraMesh3D) Vertices3() int { return len(m.Vertices) / 3 }
// Tetrahedra4 returns the tetrahedron count.
func (m *TetraMesh3D) Tetrahedra4() int { return len(m.Tetrahedra) / 4 }
// signedTetraVolume returns six times the signed volume of the
// tetrahedron (a, b, c, d): positive for the orientation the mesh
// stores, negative when the last two vertices are swapped, zero when
// the four points are coplanar.
func signedTetraVolume(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz float64) float64 {
u := [3]float64{bx - ax, by - ay, bz - az}
v := [3]float64{cx - ax, cy - ay, cz - az}
w := [3]float64{dx - ax, dy - ay, dz - az}
cross := [3]float64{v[1]*w[2] - v[2]*w[1], v[2]*w[0] - v[0]*w[2], v[0]*w[1] - v[1]*w[0]}
return u[0]*cross[0] + u[1]*cross[1] + u[2]*cross[2]
}
// BoundaryFaces returns the mesh's boundary faces as flat triples of
// vertex indices: a face belongs to the boundary when exactly one
// tetrahedron carries it. The triples are sorted lexicographically,
// so the result is a pure function of the mesh.
func (m *TetraMesh3D) BoundaryFaces() []int {
count := make(map[[3]int]int, len(m.Tetrahedra))
key := func(a, b, c int) [3]int {
if a > b {
a, b = b, a
}
if b > c {
b, c = c, b
}
if a > b {
a, b = b, a
}
return [3]int{a, b, c}
}
for t := 0; t < m.Tetrahedra4(); t++ {
a, b, c, d := int(m.Tetrahedra[4*t]), int(m.Tetrahedra[4*t+1]), int(m.Tetrahedra[4*t+2]), int(m.Tetrahedra[4*t+3])
count[key(a, b, c)]++
count[key(a, b, d)]++
count[key(a, c, d)]++
count[key(b, c, d)]++
}
sets := make([][3]int, 0, len(count))
for f, n := range count {
if n == 1 {
sets = append(sets, f)
}
}
slices.SortFunc(sets, func(x, y [3]int) int {
for k := range 3 {
if x[k] != y[k] {
return x[k] - y[k]
}
}
return 0
})
faces := make([]int, 0, 3*len(sets))
for _, f := range sets {
faces = append(faces, f[0], f[1], f[2])
}
return faces
}
// BoxTetraMesh3D builds the structured tetrahedralisation of the
// axis-aligned box [x0, x0+width] × [y0, y0+height] × [z0, z0+depth]
// with m by n by p cells, six tetrahedra per cell (the Kuhn
// subdivision along the cell diagonal, oriented positively). m, n and
// p must all be positive. The subdivision is conforming across cell
// faces, which makes the mesher the first port of call for tests and
// for boxes in general.
func BoxTetraMesh3D(x0, y0, z0, width, height, depth float64, m, n, p int) (*TetraMesh3D, error) {
const name = "BoxTetraMesh3D"
if m <= 0 || n <= 0 || p <= 0 {
return nil, base.Errf("%s: the cell counts must be positive, got %d by %d by %d", name, m, n, p)
}
// The same guard the triangle mesher applies: a non-finite extent
// or origin would lay out vertices at NaN or Inf and only surface
// mid-factorisation, far from the cause.
if !(width > 0) || !(height > 0) || !(depth > 0) ||
math.IsInf(width, 0) || math.IsInf(height, 0) || math.IsInf(depth, 0) ||
math.IsNaN(x0) || math.IsInf(x0, 0) || math.IsNaN(y0) || math.IsInf(y0, 0) || math.IsNaN(z0) || math.IsInf(z0, 0) {
return nil, base.Errf("%s: the extents must be finite and positive and the origin finite, got origin (%g, %g, %g), extents %g by %g by %g",
name, x0, y0, z0, width, height, depth)
}
vertices := make([]float64, 3*(m+1)*(n+1)*(p+1))
for k := range p + 1 {
for j := range n + 1 {
for i := range m + 1 {
v := 3 * ((k*(n+1)+j)*(m+1) + i)
vertices[v] = x0 + width*float64(i)/float64(m)
vertices[v+1] = y0 + height*float64(j)/float64(n)
vertices[v+2] = z0 + depth*float64(k)/float64(p)
}
}
}
at := func(i, j, k int) int64 { return int64((k*(n+1)+j)*(m+1) + i) }
// The six Kuhn paths from one cell corner to the opposite one,
// given as axis orders. An odd permutation reaches the far corner
// with negative orientation, so its last two vertices swap.
perms := [6][3]int{{0, 1, 2}, {0, 2, 1}, {1, 0, 2}, {1, 2, 0}, {2, 0, 1}, {2, 1, 0}}
tetrahedra := make([]int64, 0, 6*m*n*p)
for k := range p {
for j := range n {
for i := range m {
for _, pm := range perms {
// The path walks from the cell corner to the far
// corner, each vertex one axis-step beyond the
// previous one.
ox := [4]int{i, i, i, i}
oy := [4]int{j, j, j, j}
oz := [4]int{k, k, k, k}
for s := range 3 {
ox[s+1], oy[s+1], oz[s+1] = ox[s], oy[s], oz[s]
switch pm[s] {
case 0:
ox[s+1]++
case 1:
oy[s+1]++
default:
oz[s+1]++
}
}
odd := 0
for s1 := range 3 {
for s2 := s1 + 1; s2 < 3; s2++ {
if pm[s1] > pm[s2] {
odd++
}
}
}
v := [4]int64{at(ox[0], oy[0], oz[0]), at(ox[1], oy[1], oz[1]), at(ox[2], oy[2], oz[2]), at(ox[3], oy[3], oz[3])}
if odd%2 == 1 {
v[2], v[3] = v[3], v[2]
}
tetrahedra = append(tetrahedra, v[0], v[1], v[2], v[3])
}
}
}
}
return &TetraMesh3D{Vertices: vertices, Tetrahedra: tetrahedra}, nil
}
// tetraGradients returns the gradients of the four P1 basis functions
// on the tetrahedron (a, b, c, d) and its volume. The gradients are
// the columns of the inverse of the edge matrix whose rows are the
// vectors from d to a, b and c, which is the standard
// gradient-of-basis formula over the element's edge vectors.
func tetraGradients(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz float64) (g [4][3]float64, volume float64) {
// Rows of the edge matrix relative to d.
r0 := [3]float64{ax - dx, ay - dy, az - dz}
r1 := [3]float64{bx - dx, by - dy, bz - dz}
r2 := [3]float64{cx - dx, cy - dy, cz - dz}
// Cofactors of the edge matrix; the inverse is their transpose
// over the determinant, so column j of the inverse is row j of the
// cofactor matrix over det.
c00 := r1[1]*r2[2] - r1[2]*r2[1]
c01 := -(r1[0]*r2[2] - r1[2]*r2[0])
c02 := r1[0]*r2[1] - r1[1]*r2[0]
c10 := -(r0[1]*r2[2] - r0[2]*r2[1])
c11 := r0[0]*r2[2] - r0[2]*r2[0]
c12 := -(r0[0]*r2[1] - r0[1]*r2[0])
c20 := r0[1]*r1[2] - r0[2]*r1[1]
c21 := -(r0[0]*r1[2] - r0[2]*r1[0])
c22 := r0[0]*r1[1] - r0[1]*r1[0]
det := r0[0]*c00 + r0[1]*c01 + r0[2]*c02
g[0] = [3]float64{c00 / det, c01 / det, c02 / det}
g[1] = [3]float64{c10 / det, c11 / det, c12 / det}
g[2] = [3]float64{c20 / det, c21 / det, c22 / det}
for i := range 3 {
for k := range 3 {
g[3][k] -= g[i][k]
}
}
volume = math.Abs(signedTetraVolume(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz)) / 6
return g, volume
}
// tetraStiffness returns the P1 stiffness matrix of one tetrahedron:
// K[i][j] = κ·V·(∇λᵢ·∇λⱼ), the gradient-of-basis formula integrated
// over the element, where the gradients are constant on a linear
// element.
func tetraStiffness(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz, kappa float64) [4][4]float64 {
g, volume := tetraGradients(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz)
var k [4][4]float64
for i := range 4 {
for j := range 4 {
k[i][j] = kappa * volume * (g[i][0]*g[j][0] + g[i][1]*g[j][1] + g[i][2]*g[j][2])
}
}
return k
}
// FEMPoisson3DOptions carries the data SolvePoissonFEM3D needs beside
// the mesh and the source: the conductivity, the prescribed boundary
// values, and the optional flux boundary.
type FEMPoisson3DOptions struct {
// Kappa is the constant conductivity when KappaFunc is nil. It
// must be positive.
Kappa float64
// KappaFunc, when set, gives the conductivity at a point. It is
// evaluated at the tetrahedron centroids and must be positive
// there for every element; a non-positive value names the element.
KappaFunc func(x, y, z float64) float64
// DirichletNodes lists the vertices with prescribed values and
// DirichletValues the values in the same order. The nodes leave
// the system with their rows and columns; at least one is
// required, because a purely Neumann problem has no unique
// solution.
DirichletNodes []int
DirichletValues []float64
// NeumannFaces lists boundary faces as flat triples of vertex
// indices and NeumannFlux gives the flux κ∂u/∂n along each face's
// outward normal: each face's integral is built from the degree-2
// edge-midpoint rule, a third of area·flux at each edge midpoint
// shared by that edge's two vertices. A nil flux means zero.
NeumannFaces []int
NeumannFlux func(x, y, z float64) float64
// Ordering selects the fill-reducing permutation for the sparse
// Cholesky factorisation. The zero value is the natural order;
// meshes usually want SparseOrderingReverseCuthillMcKee.
Ordering linalg.SparseOrdering
}
// SolvePoissonFEM3D solves −∇·(κ∇u) = f on the tetrahedral mesh with
// piecewise-linear elements: the stiffness matrix is assembled per
// tetrahedron (the conductivity evaluated at the centroids when it
// varies), the load is integrated per tetrahedron with the 3×3×3
// collapsed Gauss rule (exact through degree 5; the centroid lump
// does not hold the O(h²) rate on the structured Kuhn mesh), Neumann
// fluxes are integrated on their boundary faces with the degree-2
// edge-midpoint rule, and Dirichlet values are eliminated by lifting.
// f may be nil for the homogeneous equation. The error contract
// mirrors SolvePoissonFEM2D.
func SolvePoissonFEM3D(mesh *TetraMesh3D, f func(x, y, z float64) float64, opts FEMPoisson3DOptions) (*core.Array, error) {
const name = "SolvePoissonFEM3D"
if mesh == nil {
return nil, base.Errf("%s: the mesh must not be nil", name)
}
// The same conductivity gate as the two-dimensional solve: with
// KappaFunc nil the constant is the value used, so it must be
// positive and finite; with the field set the constant is a
// placeholder, but a non-finite one is still refused.
if opts.KappaFunc == nil {
if !(opts.Kappa > 0) || math.IsInf(opts.Kappa, 0) {
return nil, base.Errf("%s: the conductivity must be positive, got %g", name, opts.Kappa)
}
} else if math.IsNaN(opts.Kappa) || math.IsInf(opts.Kappa, 0) {
return nil, base.Errf("%s: the conductivity must be positive, got %g", name, opts.Kappa)
}
if len(opts.DirichletNodes) != len(opts.DirichletValues) {
return nil, base.Errf("%s: %d Dirichlet nodes but %d values", name, len(opts.DirichletNodes), len(opts.DirichletValues))
}
if len(opts.DirichletNodes) == 0 {
return nil, base.Errf("%s: a purely Neumann problem has no unique solution; prescribe at least one Dirichlet value", name)
}
n := mesh.Vertices3()
// The Dirichlet nodes as a dense marker with their prescribed
// values, exactly as the two-dimensional solve carries them: the
// lifting and the unit rows each visit every assembled entry, and a
// marker answers those visits in constant time where a set of nodes
// answered with a hash. A node listed twice keeps its last value
// and appears once, as it did in the set; the appended order does
// not reach the assembled system, whose coordinate entries the
// sparse conversion sorts and merges by coordinate.
dirichletMark := make([]bool, n)
dirichletVal := make([]float64, n)
dirichletNodes := make([]int, 0, len(opts.DirichletNodes))
for p, d := range opts.DirichletNodes {
if d < 0 || d >= n {
return nil, base.Errf("%s: Dirichlet node %d out of range for %d vertices", name, d, n)
}
v := opts.DirichletValues[p]
if math.IsNaN(v) || math.IsInf(v, 0) {
return nil, base.Errf("%s: Dirichlet value at node %d is not finite", name, d)
}
if !dirichletMark[d] {
dirichletNodes = append(dirichletNodes, d)
}
dirichletMark[d] = true
dirichletVal[d] = v
}
if len(opts.NeumannFaces)%3 != 0 {
return nil, base.Errf("%s: %d Neumann face indices, want triples", name, len(opts.NeumannFaces))
}
for p := 0; p < len(opts.NeumannFaces); p += 3 {
for _, v := range opts.NeumannFaces[p : p+3] {
if v < 0 || v >= n {
return nil, base.Errf("%s: Neumann face [%d %d %d] holds the out-of-range vertex %d",
name, opts.NeumannFaces[p], opts.NeumannFaces[p+1], opts.NeumannFaces[p+2], v)
}
}
if opts.NeumannFaces[p] == opts.NeumannFaces[p+1] ||
opts.NeumannFaces[p] == opts.NeumannFaces[p+2] ||
opts.NeumannFaces[p+1] == opts.NeumannFaces[p+2] {
return nil, base.Errf("%s: Neumann face [%d %d %d] repeats a vertex",
name, opts.NeumannFaces[p], opts.NeumannFaces[p+1], opts.NeumannFaces[p+2])
}
}
// Assembly: sixteen entries per tetrahedron, symmetric by
// construction, with the conductivity evaluated at the centroid
// when it varies.
entries := make([]float64, 0, 16*mesh.Tetrahedra4())
rows := make([]int, 0, 16*mesh.Tetrahedra4())
cols := make([]int, 0, 16*mesh.Tetrahedra4())
load := make([]float64, n)
// The collapsed Gauss rule's abscissae and weights are constants of
// the scheme: built once here, not per tetrahedron.
gl := [3]float64{(1 - math.Sqrt(3.0/5)) / 2, 0.5, (1 + math.Sqrt(3.0/5)) / 2}
gw := [3]float64{5.0 / 18, 4.0 / 9, 5.0 / 18}
for t := 0; t < mesh.Tetrahedra4(); t++ {
a, b, c, d := int(mesh.Tetrahedra[4*t]), int(mesh.Tetrahedra[4*t+1]), int(mesh.Tetrahedra[4*t+2]), int(mesh.Tetrahedra[4*t+3])
ax, ay, az := mesh.Vertices[3*a], mesh.Vertices[3*a+1], mesh.Vertices[3*a+2]
bx, by, bz := mesh.Vertices[3*b], mesh.Vertices[3*b+1], mesh.Vertices[3*b+2]
cx, cy, cz := mesh.Vertices[3*c], mesh.Vertices[3*c+1], mesh.Vertices[3*c+2]
dx, dy, dz := mesh.Vertices[3*d], mesh.Vertices[3*d+1], mesh.Vertices[3*d+2]
volume := math.Abs(signedTetraVolume(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz)) / 6
if volume == 0 {
return nil, base.Errf("%s: tetrahedron %d is degenerate (zero volume)", name, t)
}
kappa := opts.Kappa
if opts.KappaFunc != nil {
kappa = opts.KappaFunc((ax+bx+cx+dx)/4, (ay+by+cy+dy)/4, (az+bz+cz+dz)/4)
if !(kappa > 0) || math.IsNaN(kappa) || math.IsInf(kappa, 0) {
return nil, base.Errf("%s: the conductivity at tetrahedron %d is %g, want positive", name, t, kappa)
}
}
k := tetraStiffness(ax, ay, az, bx, by, bz, cx, cy, cz, dx, dy, dz, kappa)
nodes := [4]int{a, b, c, d}
for i := range 4 {
for j := range 4 {
rows = append(rows, nodes[i])
cols = append(cols, nodes[j])
entries = append(entries, k[i][j])
}
}
// The load on this element, integrated with the 3×3×3
// collapsed Gauss rule: λ weights follow the Duffy collapse
// toward vertex a, and the Jacobian of the map from the unit
// cube is (1−r)²(1−s)·6V.
if f != nil {
for ir := range 3 {
for is := range 3 {
for it := range 3 {
r, s, t := gl[ir], gl[is], gl[it]
la := (1 - r) * (1 - s) * (1 - t)
lb := (1 - r) * (1 - s) * t
lc := (1 - r) * s
ld := r
x := la*ax + lb*bx + lc*cx + ld*dx
y := la*ay + lb*by + lc*cy + ld*dy
z := la*az + lb*bz + lc*cz + ld*dz
w := gw[ir] * gw[is] * gw[it] * (1 - r) * (1 - r) * (1 - s) * 6 * volume
fv := f(x, y, z)
// A non-finite source value would flow into the
// load and the solve would publish an all-NaN
// solution with a nil error, the breach every
// other integrator here refuses up front.
if math.IsNaN(fv) || math.IsInf(fv, 0) {
return nil, base.Errf("%s: the source returned the non-finite value %g at tetrahedron %d", name, fv, t)
}
load[a] += w * fv * la
load[b] += w * fv * lb
load[c] += w * fv * lc
load[d] += w * fv * ld
}
}
}
}
}
// Neumann fluxes: the degree-2 edge-midpoint rule on every listed
// face, a third of area·flux at each edge midpoint into that
// edge's two vertices.
if len(opts.NeumannFaces) > 0 && opts.NeumannFlux != nil {
for p := 0; p < len(opts.NeumannFaces); p += 3 {
a, b, c := opts.NeumannFaces[p], opts.NeumannFaces[p+1], opts.NeumannFaces[p+2]
ax, ay, az := mesh.Vertices[3*a], mesh.Vertices[3*a+1], mesh.Vertices[3*a+2]
bx, by, bz := mesh.Vertices[3*b], mesh.Vertices[3*b+1], mesh.Vertices[3*b+2]
cx, cy, cz := mesh.Vertices[3*c], mesh.Vertices[3*c+1], mesh.Vertices[3*c+2]
u := [3]float64{bx - ax, by - ay, bz - az}
v := [3]float64{cx - ax, cy - ay, cz - az}
cross := [3]float64{u[1]*v[2] - u[2]*v[1], u[2]*v[0] - u[0]*v[2], u[0]*v[1] - u[1]*v[0]}
area := math.Sqrt(cross[0]*cross[0]+cross[1]*cross[1]+cross[2]*cross[2]) / 2
w := area / 3
// A non-finite flux lands in the load like a non-finite
// source, so the same refusal answers it, naming the face.
fab := w * opts.NeumannFlux((ax+bx)/2, (ay+by)/2, (az+bz)/2)
fbc := w * opts.NeumannFlux((bx+cx)/2, (by+cy)/2, (bz+cz)/2)
fca := w * opts.NeumannFlux((cx+ax)/2, (cy+ay)/2, (cz+az)/2)
for _, fv := range []float64{fab, fbc, fca} {
if math.IsNaN(fv) || math.IsInf(fv, 0) {
return nil, base.Errf("%s: the Neumann flux returned a non-finite value on face [%d %d %d]", name, a, b, c)
}
}
load[a] += fab/2 + fca/2
load[b] += fab/2 + fbc/2
load[c] += fbc/2 + fca/2
}
}
// Dirichlet lifting: the known boundary values move to the right
// hand side, then their rows and columns leave the system as unit
// rows, exactly as in the two-dimensional solve.
for p, i := range rows {
if j := cols[p]; dirichletMark[j] {
load[i] -= entries[p] * dirichletVal[j]
}
}
keptRows := make([]int64, 0, len(rows))
keptCols := make([]int64, 0, len(rows))
keptVals := make([]float64, 0, len(rows))
for p := range rows {
i, j := rows[p], cols[p]
if dirichletMark[i] || dirichletMark[j] {
continue
}
keptRows = append(keptRows, int64(i))
keptCols = append(keptCols, int64(j))
keptVals = append(keptVals, entries[p])
}
for _, d := range dirichletNodes {
keptRows = append(keptRows, int64(d))
keptCols = append(keptCols, int64(d))
keptVals = append(keptVals, 1)
load[d] = dirichletVal[d]
}
indices, err := core.FromInts(pairInts(keptRows, keptCols), len(keptVals), 2)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
coo, err := core.NewSparseCOO(indices, fromSlice(keptVals, len(keptVals)), []int{n, n})
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
factor, err := linalg.NewSparseCholesky(coo, opts.Ordering)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
rhs := core.New(core.Float, []int{n}...)
copy(rhs.RawFloats(), load)
return factor.Solve(rhs)
}
+565
View File
@@ -0,0 +1,565 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"testing"
linalg "sourcedock.dev/petrbalvin/tensor/linalg"
)
// boxMesh3D builds the structured tetrahedralisation of the unit box
// with m cells per side and returns the mesh plus the list of
// boundary vertices, in mesh order.
func boxMesh3D(t *testing.T, m int) (*TetraMesh3D, []int) {
t.Helper()
mesh, err := BoxTetraMesh3D(0, 0, 0, 1, 1, 1, m, m, m)
if err != nil {
t.Fatalf("BoxTetraMesh3D: %v", err)
}
boundary := make([]int, 0, 6*(m+1)*(m+1))
for k := range m + 1 {
for j := range m + 1 {
for i := range m + 1 {
if i == 0 || i == m || j == 0 || j == m || k == 0 || k == m {
boundary = append(boundary, (k*(m+1)+j)*(m+1)+i)
}
}
}
}
return mesh, boundary
}
// TestTetraStiffnessReference pins the element stiffness matrix
// against the hand-computed 4x4 for the reference tetrahedron
// (0,0,0), (1,0,0), (0,1,0), (0,0,1): with κ = 1 the matrix is
// κ/6·[[3,−1,−1,−1],[−1,1,0,0],[−1,0,1,0],[−1,0,0,1]].
func TestTetraStiffnessReference(t *testing.T) {
hand := [4][4]float64{
{3, -1, -1, -1},
{-1, 1, 0, 0},
{-1, 0, 1, 0},
{-1, 0, 0, 1},
}
k := tetraStiffness(0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 1)
for i := range 4 {
for j := range 4 {
want := hand[i][j] / 6
if math.Abs(k[i][j]-want) > 1e-15 {
t.Fatalf("K[%d][%d] = %.17g, want %.17g", i, j, k[i][j], want)
}
}
}
// The conductivity scales the matrix, nothing else.
k2 := tetraStiffness(0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 2.5)
for i := range 4 {
for j := range 4 {
if math.Abs(k2[i][j]-2.5*k[i][j]) > 1e-15 {
t.Fatalf("K[%d][%d] did not scale with κ", i, j)
}
}
}
// A tetrahedron scaled by two in every direction: the basis
// gradients halve and the volume grows eightfold, so each entry
// doubles.
ks := tetraStiffness(0, 0, 0, 2, 0, 0, 0, 2, 0, 0, 0, 2, 1)
for i := range 4 {
for j := range 4 {
if math.Abs(ks[i][j]-2*k[i][j]) > 1e-14 {
t.Fatalf("scaled K[%d][%d] = %.17g, want %.17g", i, j, ks[i][j], 2*k[i][j])
}
}
}
// The matrix is symmetric with positive diagonals and zero row
// sums off the constant mode: the P1 rigid-body mode has no
// stiffness.
for i := range 4 {
sum := 0.0
for j := range 4 {
if math.Abs(k[i][j]-k[j][i]) > 1e-15 {
t.Fatalf("K[%d][%d] != K[%d][%d]", i, j, j, i)
}
sum += k[i][j]
}
if math.Abs(sum) > 1e-14 {
t.Fatalf("row %d sums to %.3g, want 0", i, sum)
}
}
}
// TestBoxTetraMesh3DStructure pins the structured mesher: the vertex
// and tetrahedron counts, exact corner coordinates, positive
// orientation everywhere, unit total volume, and the boundary face
// count of the box surface.
func TestBoxTetraMesh3DStructure(t *testing.T) {
m, n, p := 3, 2, 4
mesh, err := BoxTetraMesh3D(0.5, -1, 2, 1.5, 1, 2, m, n, p)
if err != nil {
t.Fatalf("BoxTetraMesh3D: %v", err)
}
if mesh.Vertices3() != (m+1)*(n+1)*(p+1) {
t.Fatalf("vertex count %d, want %d", mesh.Vertices3(), (m+1)*(n+1)*(p+1))
}
if mesh.Tetrahedra4() != 6*m*n*p {
t.Fatalf("tetrahedron count %d, want %d", mesh.Tetrahedra4(), 6*m*n*p)
}
// Exact corner coordinates of the box.
at := func(i, j, k int) int { return (k*(n+1)+j)*(m+1) + i }
checkCorner := func(label string, i, j, k int, want [3]float64) {
t.Helper()
v := 3 * at(i, j, k)
for d := range 3 {
if mesh.Vertices[v+d] != want[d] {
t.Fatalf("%s = (%g, %g, %g), want (%g, %g, %g)",
label, mesh.Vertices[v], mesh.Vertices[v+1], mesh.Vertices[v+2], want[0], want[1], want[2])
}
}
}
checkCorner("origin", 0, 0, 0, [3]float64{0.5, -1, 2})
checkCorner("far corner", m, n, p, [3]float64{2, 0, 4})
// Every tetrahedron positively oriented, and the volumes sum to
// the box volume: 1.5 · 1 · 2 = 3.
total := 0.0
cell := 3.0 / float64(6*m*n*p)
for t4 := range mesh.Tetrahedra4() {
a, b, c, d := int(mesh.Tetrahedra[4*t4]), int(mesh.Tetrahedra[4*t4+1]), int(mesh.Tetrahedra[4*t4+2]), int(mesh.Tetrahedra[4*t4+3])
s6 := signedTetraVolume(
mesh.Vertices[3*a], mesh.Vertices[3*a+1], mesh.Vertices[3*a+2],
mesh.Vertices[3*b], mesh.Vertices[3*b+1], mesh.Vertices[3*b+2],
mesh.Vertices[3*c], mesh.Vertices[3*c+1], mesh.Vertices[3*c+2],
mesh.Vertices[3*d], mesh.Vertices[3*d+1], mesh.Vertices[3*d+2])
if s6 <= 0 {
t.Fatalf("tetrahedron %d has signed volume %g", t4, s6/6)
}
if d := math.Abs(s6/6 - cell); d > 1e-12 {
t.Fatalf("tetrahedron %d has volume %.6g, want %.6g", t4, s6/6, cell)
}
total += s6 / 6
}
if math.Abs(total-3) > 1e-12 {
t.Fatalf("total volume %.6g, want 3", total)
}
// The box surface carries two triangles per unit square face.
faces := mesh.BoundaryFaces()
if len(faces) != 3*2*2*(m*n+n*p+m*p) {
t.Fatalf("boundary face triples %d, want %d", len(faces)/3, 2*2*(m*n+n*p+m*p))
}
// Every listed face holds three distinct vertices, all on the box
// surface.
onSurface := func(v int) bool {
i := v % (m + 1)
j := (v / (m + 1)) % (n + 1)
k := v / ((m + 1) * (n + 1))
return i == 0 || i == m || j == 0 || j == n || k == 0 || k == p
}
for q := 0; q < len(faces); q += 3 {
if faces[q] == faces[q+1] || faces[q] == faces[q+2] || faces[q+1] == faces[q+2] {
t.Fatalf("boundary face [%d %d %d] repeats a vertex", faces[q], faces[q+1], faces[q+2])
}
for r := range 3 {
if !onSurface(faces[q+r]) {
t.Fatalf("boundary face vertex %d is interior", faces[q+r])
}
}
}
}
func TestBoxTetraMesh3DRefusals(t *testing.T) {
if _, err := BoxTetraMesh3D(0, 0, 0, 1, 1, 1, 0, 2, 2); err == nil {
t.Fatal("a zero cell count was accepted")
}
if _, err := BoxTetraMesh3D(0, 0, 0, 1, -1, 1, 2, 2, 2); err == nil {
t.Fatal("a negative extent was accepted")
}
if _, err := BoxTetraMesh3D(math.NaN(), 0, 0, 1, 1, 1, 2, 2, 2); err == nil {
t.Fatal("a NaN origin was accepted")
}
if _, err := BoxTetraMesh3D(0, 0, 0, math.Inf(1), 1, 1, 2, 2, 2); err == nil {
t.Fatal("an infinite extent was accepted")
}
}
// TestTetraMesh3DRefusals pins the construction contract: shapes,
// dtypes, ranges, finiteness, and the refusal of degenerate and
// inverted tetrahedra with the offending coordinates named.
func TestTetraMesh3DRefusals(t *testing.T) {
good := []float64{0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1}
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, good, 4, 3), floatsToArrayFEM(t, []float64{0, 1, 2, 3}, 1, 4)); err == nil || !stringsContains(err, "integers") {
t.Fatalf("a float tetrahedron table: %v", err)
}
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, good[:9], 3, 3), intsToArrayFEM(t, []int64{0, 1, 2, 3}, 1, 4)); err == nil || !stringsContains(err, "at least four vertices") {
t.Fatalf("three vertices: %v", err)
}
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, good, 4, 3), intsToArrayFEM(t, []int64{}, 0, 4)); err == nil || !stringsContains(err, "must not be empty") {
t.Fatalf("an empty tetrahedron table: %v", err)
}
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, good, 4, 3), intsToArrayFEM(t, []int64{0, 1, 2, 9}, 1, 4)); err == nil || !stringsContains(err, "out of range") {
t.Fatalf("an out-of-range index: %v", err)
}
bad := append([]float64{}, good...)
bad[0] = math.NaN()
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, bad, 4, 3), intsToArrayFEM(t, []int64{0, 1, 2, 3}, 1, 4)); err == nil || !stringsContains(err, "not finite") {
t.Fatalf("a NaN coordinate: %v", err)
}
// Degenerate: four coplanar points.
degenerate := []float64{0, 0, 0, 1, 0, 0, 0, 1, 0, 1, 1, 0}
_, err := NewTetraMesh3D(floatsToArrayFEM(t, degenerate, 4, 3), intsToArrayFEM(t, []int64{0, 1, 2, 3}, 1, 4))
if err == nil || !stringsContains(err, "degenerate") {
t.Fatalf("a coplanar tetrahedron: %v", err)
}
// Inverted: the reference tetrahedron with its last two vertices
// swapped; the message names the coordinates.
inverted := []float64{0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 1, 0}
_, err = NewTetraMesh3D(floatsToArrayFEM(t, inverted, 4, 3), intsToArrayFEM(t, []int64{0, 1, 2, 3}, 1, 4))
if err == nil || !stringsContains(err, "inverted") {
t.Fatalf("an inverted tetrahedron: %v", err)
}
if indexOf(err.Error(), "-0.1666") < 0 {
t.Fatalf("the inverted message should name the negative signed volume: %v", err)
}
if !stringsContains(err, "(1, 0, 0)") {
t.Fatalf("the inverted message should name the coordinates: %v", err)
}
// Wrong vertex table shape.
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, good[:8], 4, 2), intsToArrayFEM(t, []int64{0, 1, 2, 3}, 1, 4)); err == nil || !stringsContains(err, "three columns") {
t.Fatalf("a two-column vertex table: %v", err)
}
if _, err := NewTetraMesh3D(floatsToArrayFEM(t, good, 4, 3), intsToArrayFEM(t, []int64{0, 1, 2}, 1, 3)); err == nil || !stringsContains(err, "four columns") {
t.Fatalf("a three-column tetrahedron table: %v", err)
}
}
// TestSolvePoissonFEM3DConvergence runs the manufactured solution
// u = sin(πx)·sin(πy)·sin(πz) on the unit box, driven by
// f = 3π²·u: the P1 nodal error must keep the O(h²) rate, roughly
// quadrupling per mesh doubling, exactly as the two-dimensional solve
// pins.
func TestSolvePoissonFEM3DConvergence(t *testing.T) {
solution := func(x, y, z float64) float64 {
return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) * math.Sin(math.Pi*z)
}
source := func(x, y, z float64) float64 { return 3 * math.Pi * math.Pi * solution(x, y, z) }
previous := 0.0
for _, m := range []int{4, 8, 16} {
mesh, boundary := boxMesh3D(t, m)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[3*node], mesh.Vertices[3*node+1], mesh.Vertices[3*node+2])
}
u, err := SolvePoissonFEM3D(mesh, source, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM3D(m=%d): %v", m, err)
}
worst := 0.0
for i := range mesh.Vertices3() {
d := math.Abs(u.FloatAt(i) - solution(mesh.Vertices[3*i], mesh.Vertices[3*i+1], mesh.Vertices[3*i+2]))
if d > worst {
worst = d
}
}
t.Logf("m=%2d: max nodal error %.3g", m, worst)
if previous > 0 && previous/worst < 2.5 {
t.Fatalf("m=%d: refinement ratio %.2f, want the O(h²) rate (previous %.3g, now %.3g)",
m, previous/worst, previous, worst)
}
previous = worst
}
}
// TestSolvePoissonFEM3DVariableKappa mirrors the two-dimensional
// variable-conductivity pin on the axis a centroid typo once
// corrupted: with κ = 1 + y the conductivity sample each element sees
// comes from its own y centroid, and the P1 convergence rate must
// survive the varying coefficient.
func TestSolvePoissonFEM3DVariableKappa(t *testing.T) {
solution := func(x, y, z float64) float64 {
return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) * math.Sin(math.Pi*z)
}
kappaF := func(x, y, z float64) float64 { return 1 + y }
uy := func(x, y, z float64) float64 {
return math.Pi * math.Sin(math.Pi*x) * math.Cos(math.Pi*y) * math.Sin(math.Pi*z)
}
source := func(x, y, z float64) float64 {
return 3*math.Pi*math.Pi*(1+y)*solution(x, y, z) - uy(x, y, z)
}
previous := 0.0
for _, m := range []int{4, 8, 16} {
mesh, boundary := boxMesh3D(t, m)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[3*node], mesh.Vertices[3*node+1], mesh.Vertices[3*node+2])
}
u, err := SolvePoissonFEM3D(mesh, source, FEMPoisson3DOptions{KappaFunc: kappaF, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM3D(m=%d): %v", m, err)
}
worst := 0.0
for i := range mesh.Vertices3() {
d := math.Abs(u.FloatAt(i) - solution(mesh.Vertices[3*i], mesh.Vertices[3*i+1], mesh.Vertices[3*i+2]))
if d > worst {
worst = d
}
}
t.Logf("m=%2d: max nodal error %.3g", m, worst)
if previous > 0 && previous/worst < 2.5 {
t.Fatalf("m=%d: refinement ratio %.2f, want the O(h²) rate (previous %.3g, now %.3g)",
m, previous/worst, previous, worst)
}
previous = worst
}
}
// TestSolvePoissonFEM3DPatchLinear is the patch test: a linear field
// lies in the P1 space, so with f = 0 and the boundary lifted the
// interior solution must equal the field to machine precision.
func TestSolvePoissonFEM3DPatchLinear(t *testing.T) {
mesh, boundary := boxMesh3D(t, 6)
field := func(x, y, z float64) float64 { return 1 + 2*x - 3*y + 4*z }
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = field(mesh.Vertices[3*node], mesh.Vertices[3*node+1], mesh.Vertices[3*node+2])
}
u, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM3D: %v", err)
}
worst := 0.0
for i := range mesh.Vertices3() {
d := math.Abs(u.FloatAt(i) - field(mesh.Vertices[3*i], mesh.Vertices[3*i+1], mesh.Vertices[3*i+2]))
if d > worst {
worst = d
}
}
if worst > 1e-11 {
t.Fatalf("linear patch test error %.3g, want machine precision", worst)
}
}
// TestSolvePoissonFEM3DNeumannNatural pins the natural boundary: a
// constant field with f = 0 satisfies the homogeneous Neumann
// condition everywhere, so pinning the constant at a single vertex
// must reproduce it across the whole mesh.
func TestSolvePoissonFEM3DNeumannNatural(t *testing.T) {
mesh, _ := boxMesh3D(t, 5)
u, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{4}})
if err != nil {
t.Fatalf("SolvePoissonFEM3D: %v", err)
}
for i := range mesh.Vertices3() {
if math.Abs(u.FloatAt(i)-4) > 1e-10 {
t.Fatalf("node %d: solution %.12g, want the constant 4", i, u.FloatAt(i))
}
}
}
// TestSolvePoissonFEM3DNeumannFlux pins the boundary-face integrals:
// u = (x²+y²+z²)/2 has −Δu = −3 and the flux κ∂u/∂n = 1 on the three
// faces at x = 1, y = 1 and z = 1 (0 on the coordinate planes), so
// prescribing those fluxes with a single pinned vertex must
// reproduce the quadratic field to the accuracy of the edge-midpoint
// face rule, improving as the mesh refines.
func TestSolvePoissonFEM3DNeumannFlux(t *testing.T) {
field := func(x, y, z float64) float64 { return (x*x + y*y + z*z) / 2 }
previous := 0.0
for _, m := range []int{8, 16} {
mesh, err := BoxTetraMesh3D(0, 0, 0, 1, 1, 1, m, m, m)
if err != nil {
t.Fatalf("BoxTetraMesh3D: %v", err)
}
var faces []int
bf := mesh.BoundaryFaces()
for p := 0; p < len(bf); p += 3 {
f := bf[p : p+3]
mx := (mesh.Vertices[3*f[0]] + mesh.Vertices[3*f[1]] + mesh.Vertices[3*f[2]]) / 3
my := (mesh.Vertices[3*f[0]+1] + mesh.Vertices[3*f[1]+1] + mesh.Vertices[3*f[2]+1]) / 3
mz := (mesh.Vertices[3*f[0]+2] + mesh.Vertices[3*f[1]+2] + mesh.Vertices[3*f[2]+2]) / 3
if mx == 1 || my == 1 || mz == 1 {
faces = append(faces, f[0], f[1], f[2])
}
}
flux := func(x, y, z float64) float64 {
if x == 1 || y == 1 || z == 1 {
return 1
}
return 0
}
u, err := SolvePoissonFEM3D(mesh, func(float64, float64, float64) float64 { return -3 },
FEMPoisson3DOptions{
Kappa: 1,
DirichletNodes: []int{0},
DirichletValues: []float64{0},
NeumannFaces: faces,
NeumannFlux: flux,
})
if err != nil {
t.Fatalf("m=%d: %v", m, err)
}
worst := 0.0
for i := range mesh.Vertices3() {
d := math.Abs(u.FloatAt(i) - field(mesh.Vertices[3*i], mesh.Vertices[3*i+1], mesh.Vertices[3*i+2]))
if d > worst {
worst = d
}
}
t.Logf("m=%2d: max nodal error %.3g", m, worst)
if previous > 0 && previous/worst < 1.3 {
t.Fatalf("m=%d: refinement ratio %.2f, want the face-rule error to shrink under refinement", m, previous/worst)
}
previous = worst
}
}
// TestSolvePoissonFEM3DOrderings runs the manufactured-solution solve
// under every ordering the factor offers: the ordering changes the
// fill, never the answer.
func TestSolvePoissonFEM3DOrderings(t *testing.T) {
solution := func(x, y, z float64) float64 {
return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) * math.Sin(math.Pi*z)
}
mesh, boundary := boxMesh3D(t, 5)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[3*node], mesh.Vertices[3*node+1], mesh.Vertices[3*node+2])
}
source := func(x, y, z float64) float64 { return 3 * math.Pi * math.Pi * solution(x, y, z) }
reference, err := SolvePoissonFEM3D(mesh, source, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values})
if err != nil {
t.Fatalf("SolvePoissonFEM3D(natural): %v", err)
}
for _, ordering := range []linalg.SparseOrdering{
linalg.SparseOrderingReverseCuthillMcKee,
linalg.SparseOrderingMinimumDegree,
} {
u, err := SolvePoissonFEM3D(mesh, source, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: values, Ordering: ordering})
if err != nil {
t.Fatalf("SolvePoissonFEM3D(%d): %v", ordering, err)
}
for i := range mesh.Vertices3() {
if math.Abs(u.FloatAt(i)-reference.FloatAt(i)) > 1e-9 {
t.Fatalf("ordering %d: node %d differs from the natural run", ordering, i)
}
}
}
}
func TestSolvePoissonFEM3DRefusals(t *testing.T) {
mesh, boundary := boxMesh3D(t, 3)
zero := make([]float64, len(boundary))
if _, err := SolvePoissonFEM3D(nil, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}}); err == nil || !stringsContains(err, "nil") {
t.Fatalf("a nil mesh: %v", err)
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 0, DirichletNodes: boundary, DirichletValues: zero}); err == nil || !stringsContains(err, "positive") {
t.Fatalf("zero conductivity: %v", err)
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0, 1}, DirichletValues: []float64{1}}); err == nil {
t.Fatal("a Dirichlet length mismatch was accepted")
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1}); err == nil || !stringsContains(err, "purely Neumann") {
t.Fatalf("a purely Neumann problem: %v", err)
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{99}, DirichletValues: []float64{1}}); err == nil || !stringsContains(err, "out of range") {
t.Fatalf("an out-of-range Dirichlet node: %v", err)
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{math.NaN()}}); err == nil || !stringsContains(err, "not finite") {
t.Fatalf("a NaN Dirichlet value: %v", err)
}
// Neumann face tables: not triples, out-of-range and repeated
// vertices.
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}, NeumannFaces: []int{0, 1, 2, 3}}); err == nil || !stringsContains(err, "triples") {
t.Fatalf("a Neumann face count not divisible by three: %v", err)
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}, NeumannFaces: []int{0, 1, 77}}); err == nil || !stringsContains(err, "out-of-range") {
t.Fatalf("an out-of-range Neumann vertex: %v", err)
}
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}, NeumannFaces: []int{0, 0, 1}}); err == nil || !stringsContains(err, "repeats") {
t.Fatalf("a degenerate Neumann face: %v", err)
}
// A KappaFunc returning a non-positive conductivity names the
// tetrahedron.
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{
KappaFunc: func(float64, float64, float64) float64 { return -1 },
DirichletNodes: []int{0},
DirichletValues: []float64{0},
}); err == nil || !stringsContains(err, "positive") {
t.Fatalf("a non-positive KappaFunc value: %v", err)
}
// A KappaFunc returning an infinite conductivity names the
// tetrahedron the way the constant field's gate names itself.
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{
KappaFunc: func(float64, float64, float64) float64 { return math.Inf(1) },
DirichletNodes: []int{0},
DirichletValues: []float64{0},
}); err == nil || !stringsContains(err, "positive") {
t.Fatalf("an infinite KappaFunc value: %v", err)
}
// A non-finite source value refuses the solve: it used to land in
// the load and publish an all-NaN solution with a nil error.
if _, err := SolvePoissonFEM3D(mesh, func(x, y, z float64) float64 { return math.NaN() },
FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: zero}); err == nil || !stringsContains(err, "non-finite") {
t.Fatalf("a NaN source value: %v", err)
}
// A non-finite Neumann flux refuses the solve the same way.
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{
Kappa: 1,
DirichletNodes: []int{0},
DirichletValues: []float64{0},
NeumannFaces: []int{0, 1, 2},
NeumannFlux: func(x, y, z float64) float64 { return math.Inf(1) },
}); err == nil || !stringsContains(err, "non-finite") {
t.Fatalf("an infinite Neumann flux: %v", err)
}
// An ordering that does not exist.
if _, err := SolvePoissonFEM3D(mesh, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: zero, Ordering: linalg.SparseOrdering(7)}); err == nil {
t.Fatal("an unknown ordering was accepted")
}
// A hand-built mesh with a degenerate tetrahedron is refused by
// the solver, which checks the volume itself.
hollow := &TetraMesh3D{
Vertices: []float64{0, 0, 0, 1, 0, 0, 0, 1, 0, 1, 1, 0},
Tetrahedra: []int64{0, 1, 2, 3},
}
if _, err := SolvePoissonFEM3D(hollow, nil, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: []int{0}, DirichletValues: []float64{0}}); err == nil || !stringsContains(err, "degenerate") {
t.Fatalf("a hand-built degenerate mesh: %v", err)
}
}
// TestSolvePoissonFEM3DDuplicateDirichletNode pins the three-dimensional
// side of the same rule as the two-dimensional test: a node listed
// twice keeps its last value and is recorded once, so the repeated
// listing answers what the single listing with that value answers.
func TestSolvePoissonFEM3DDuplicateDirichletNode(t *testing.T) {
solution := func(x, y, z float64) float64 {
return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) * math.Sin(math.Pi*z)
}
source := func(x, y, z float64) float64 { return 3 * math.Pi * math.Pi * solution(x, y, z) }
mesh, boundary := boxMesh3D(t, 2)
values := make([]float64, len(boundary))
for p, node := range boundary {
values[p] = solution(mesh.Vertices[3*node], mesh.Vertices[3*node+1], mesh.Vertices[3*node+2])
}
const extra = 0.5
nodes := append(append([]int(nil), boundary...), boundary[1])
dupValues := append(append([]float64(nil), values...), values[1]+extra)
u, err := SolvePoissonFEM3D(mesh, source, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: nodes, DirichletValues: dupValues})
if err != nil {
t.Fatalf("SolvePoissonFEM3D with a repeated node: %v", err)
}
single := append([]float64(nil), values...)
single[1] += extra
want, err := SolvePoissonFEM3D(mesh, source, FEMPoisson3DOptions{Kappa: 1, DirichletNodes: boundary, DirichletValues: single})
if err != nil {
t.Fatalf("SolvePoissonFEM3D with the node once: %v", err)
}
if got := u.FloatAt(boundary[1]); math.Abs(got-(values[1]+extra)) > 1e-12 {
t.Fatalf("the repeated node answered %g, want the last prescribed value %g", got, values[1]+extra)
}
worst := 0.0
for i := range mesh.Vertices3() {
worst = math.Max(worst, math.Abs(u.FloatAt(i)-want.FloatAt(i)))
}
if worst > 1e-12 {
t.Fatalf("the repeated listing differs from the single listing by %g, want the node recorded once", worst)
}
}
+234
View File
@@ -0,0 +1,234 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// Oscillatory quadrature: the integral of a smooth amplitude against a
// sine or cosine of a high frequency, the shape every spectral
// reduction produces and one a plain adaptive rule pays for double: it
// must resolve the carrier, not the amplitude, so the evaluation count
// grows with the frequency and the per-panel rules start aliasing.
//
// The scheme is Filon-type. The interval splits into equal panels, the
// amplitude f is interpolated on each panel by a polynomial through
// Gauss-Legendre nodes, and the product of that polynomial with the
// oscillatory kernel is carried out exactly through per-panel weights.
// The error therefore tracks the smoothness of f alone and falls like
// the panel width to the interpolation order, no matter how large the
// frequency grows, while the plain adaptive rule must spend roughly
// twenty evaluations per carrier wavelength to see it at all.
// FilonOptions tunes IntegrateFilon. Nodes ≤ 0 means 16, the
// polynomial degree of the amplitude interpolant per panel is Nodes−1.
// Panels ≤ 0 means automatic: the count that keeps each panel at most
// about Nodes half-wavelengths of the carrier, the range where the
// moment construction below is exact to the rounding floor.
type FilonOptions struct {
Panels int
Nodes int
}
// filonAlphaCap bounds the forced-panel moment phase: a panel may
// carry at most this many half-wavelengths of the carrier before the
// auxiliary rule that builds the weights would have to grow without
// bound. The automatic panel count never reaches it.
const filonAlphaCap = 4096.0
// IntegrateFilon returns the two definite integrals
//
// cosIntegral = ∫ f(x)·cos(kx) dx, sinIntegral = ∫ f(x)·sin(kx) dx
//
// over [a, b], the real and imaginary parts of ∫ f(x)·e^{ikx} dx. A
// reversed interval integrates in the negative direction and k = 0
// degenerates to the plain integral of f with a zero sine part. The
// construction is exact whenever f is a polynomial of degree below
// Nodes, so on smooth amplitudes the answer sits at the rounding floor
// even for frequencies whose carrier a sampled rule cannot see.
//
// Errors: NaN or infinite bounds, an infinite frequency, a NaN
// frequency, Nodes outside [2, 32], a forced Panels whose panels would
// carry more than filonAlphaCap half-wavelengths of the carrier, a span
// that overflows the float64 range, a frequency whose span product
// leaves no representable panel count, and an f that fails or returns
// a non-finite value.
func IntegrateFilon(f func(x float64) (float64, error), a, b, k float64, opts FilonOptions) (float64, float64, error) {
if opts.Nodes <= 0 {
opts.Nodes = 16
}
if opts.Nodes < 2 || opts.Nodes > 32 {
return 0, 0, base.Errf("IntegrateFilon: Nodes must be between 2 and 32, got %d", opts.Nodes)
}
if math.IsNaN(a) || math.IsNaN(b) || math.IsNaN(k) {
return 0, 0, base.Errf("IntegrateFilon: bounds and frequency must not be NaN")
}
if math.IsInf(a, 0) || math.IsInf(b, 0) {
return 0, 0, base.Errf("IntegrateFilon: bounds must be finite, got [%g, %g]", a, b)
}
if math.IsInf(k, 0) {
return 0, 0, base.Errf("IntegrateFilon: the frequency must be finite, got %g", k)
}
sign := 1.0
if b < a {
a, b = b, a
sign = -1
}
if a == b {
return 0, 0, nil
}
// Two finite bounds can still sit so far apart that their span
// overflows: the panel width would be infinite and the carrier's
// phase at the panel centre 0·Inf or k·Inf, a quiet NaN pair.
if span := b - a; math.IsInf(span, 0) {
return 0, 0, base.Errf("IntegrateFilon: the span from %g to %g overflows, leaving no representable panel width", a, b)
}
if opts.Panels > 0 {
if alpha := math.Abs(k) * (b - a) / (2 * float64(opts.Panels)); alpha > filonAlphaCap {
return 0, 0, base.Errf("IntegrateFilon: %d panels leave %g half-wavelengths of the carrier per panel, above the %g the weights can be built within; raise Panels or leave them automatic",
opts.Panels, alpha, filonAlphaCap)
}
}
panels := opts.Panels
if panels <= 0 {
panels = 1
if k != 0 {
// A panel of h carries |k|h/2 half-wavelengths; the cap at
// Nodes keeps the moment construction in its exact range
// and the interpolation error far under the floor. The
// estimate can also leave the int range while still
// finite, and the conversion of such a ceiling is
// implementation-dependent garbage: on saturation it asks
// for an unending loop, elsewhere it wraps negative and
// the empty loop reports a quiet zero. Refuse anything
// the platform's int cannot represent.
est := math.Abs(k) * (b - a) / (2 * float64(opts.Nodes))
if est >= math.MaxInt {
return 0, 0, base.Errf("IntegrateFilon: the frequency %g over the span %g leaves no representable panel count", k, b-a)
}
panels = int(math.Ceil(est))
}
}
h := (b - a) / float64(panels)
alpha := k * h / 2
nodes, _, err := GaussLegendreNodes(opts.Nodes)
if err != nil {
return 0, 0, err
}
wCos, wSin, err := filonWeights(nodes, alpha)
if err != nil {
return 0, 0, err
}
// One sweep over the panels: sample the amplitude at the nodes,
// contract it with the weights into the panel's two amplitudes C
// and S, and rotate them into place by the carrier's phase at the
// panel centre.
var cosTotal, sinTotal float64
for p := range panels {
centre := a + (float64(p)+0.5)*h
half := h / 2
var c, s float64
for i := range nodes {
fx, ferr := f(centre + half*nodes[i])
if ferr != nil {
return 0, 0, base.Errf("IntegrateFilon: %w", ferr)
}
if math.IsNaN(fx) || math.IsInf(fx, 0) {
return 0, 0, base.Errf("IntegrateFilon: the amplitude returned the non-finite value %g on panel %d", fx, p)
}
c += wCos[i] * fx
s += wSin[i] * fx
}
phase := k * centre
cosP, sinP := math.Cos(phase), math.Sin(phase)
cosTotal += cosP*c - sinP*s
sinTotal += sinP*c + cosP*s
}
return sign * cosTotal * h / 2, sign * sinTotal * h / 2, nil
}
// filonWeights returns, for the Gauss-Legendre nodes of [-1, 1], the
// Filon weights: the exact integrals of each Lagrange basis polynomial
// against cos(αy) and sin(αy). With these the panel integral of the
// interpolating polynomial times the carrier is one dot product per
// part, and every trace of the carrier's phase lives in the weights,
// built once, never per panel.
//
// The basis moments come from a composite 32-point Gauss-Legendre rule
// whose subinterval count follows α, so the auxiliary rule resolves
// the carrier the amplitude is multiplied by; the automatic panel cap
// keeps that cost at one subinterval and the rule at the rounding
// floor.
func filonWeights(nodes []float64, alpha float64) (wCos, wSin []float64, err error) {
m := len(nodes)
// Barycentric weights of the interpolation nodes.
bw := make([]float64, m)
for i := range m {
p := 1.0
for j := range m {
if j != i {
p *= nodes[i] - nodes[j]
}
}
if p == 0 {
return nil, nil, base.Errf("IntegrateFilon: repeated interpolation nodes")
}
bw[i] = 1 / p
}
// The auxiliary rule: 32-point Gauss-Legendre over enough equal
// subintervals of [-1, 1] that each carries at most 16
// half-wavelengths of e^{iαy}.
subs := 1
if a := math.Abs(alpha); a > 16 {
subs = int(math.Ceil(a / 16))
}
auxNodes, auxWeights, err := GaussLegendreNodes(32)
if err != nil {
return nil, nil, err
}
wCos = make([]float64, m)
wSin = make([]float64, m)
span := 2.0 / float64(subs)
for s := range subs {
lo := -1 + float64(s)*span
for t := range auxNodes {
// The aux nodes live on [-1, 1]; map them into the
// subinterval [lo, lo+span] with the half-span the affine
// change of variables carries.
y := lo + span*0.5*(auxNodes[t]+1)
// Barycentric evaluation of every basis polynomial at y,
// with the exact hit a node coincidence asks for.
den := 0.0
exact := -1
for i := range m {
d := y - nodes[i]
if d == 0 {
exact = i
break
}
den += bw[i] / d
}
cy, sy := math.Cos(alpha*y), math.Sin(alpha*y)
w := span * 0.5 * auxWeights[t]
for i := range m {
var li float64
if exact >= 0 {
if i == exact {
li = 1
}
} else {
li = bw[i] / (y - nodes[i]) / den
}
wCos[i] += w * li * cy
wSin[i] += w * li * sy
}
}
}
return wCos, wSin, nil
}
+330
View File
@@ -0,0 +1,330 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"math/big"
"slices"
"testing"
)
// IntegrateFilon against an external exact reference: the antiderivative
//
// ∫ p(x)·e^{ikx} dx = e^{ikx}·Σ_{j≥0} (−1)^j p^{(j)}(x)/(ik)^{j+1},
//
// summed in math/big at a working size far past the float64 grid, with
// π from Machin's formula and the endpoint phases reduced mod 2π before
// the Taylor run. The reference holds for every frequency tried here,
// so a phase defect of the scheme itself shows against it.
const filonRefPrec = 512
func fb(x float64) *big.Float {
return new(big.Float).SetPrec(filonRefPrec).SetFloat64(x)
}
func fbInt(n int64) *big.Float {
return new(big.Float).SetPrec(filonRefPrec).SetInt64(n)
}
func fbPi() *big.Float {
// π = 16·atan(1/5) − 4·atan(1/239).
atan := func(t *big.Float) *big.Float {
power := new(big.Float).SetPrec(filonRefPrec).Set(t)
sum := fb(0)
for k := int64(1); ; k += 2 {
term := new(big.Float).SetPrec(filonRefPrec).Quo(power, fbInt(k))
if (k/2)%2 == 1 {
term.Neg(term)
}
sum.Add(sum, term)
power.Mul(power, t)
power.Mul(power, t)
if term.MantExp(nil) < -int(filonRefPrec)-10 {
break
}
}
return sum
}
// 1/5 must reach atan as the exact quotient: the float64 literal
// 0.2 carries a 1e-17 argument error that Machin's formula
// amplifies sixteenfold into π itself.
fifth := new(big.Float).SetPrec(filonRefPrec).Quo(fb(1), fbInt(5))
two39 := new(big.Float).SetPrec(filonRefPrec).Quo(fb(1), fbInt(239))
sixteen := new(big.Float).SetPrec(filonRefPrec).Mul(fbInt(16), atan(fifth))
four := new(big.Float).SetPrec(filonRefPrec).Mul(fbInt(4), atan(two39))
return sixteen.Sub(sixteen, four)
}
var (
filonTwoPi = new(big.Float).SetPrec(filonRefPrec).Mul(fb(2), fbPi())
filonPi = fbPi()
)
// filonBigSinCos returns sin(x), cos(x) for the exact big.Float argument,
// kept in extended precision: the endpoint products below multiply them
// by antiderivative terms far larger than the integral itself, so a
// float64 detour here would show up in the reference's own answer.
func filonBigSinCos(x *big.Float) (s, c *big.Float) {
n := new(big.Float).SetPrec(filonRefPrec).Quo(x, filonTwoPi)
ni, _ := n.Int(nil)
r := new(big.Float).SetPrec(filonRefPrec).Mul(new(big.Float).SetInt(ni), filonTwoPi)
r.Sub(x, r)
// The remainder sits within (−2π, 2π); one step puts it in (−π, π].
halfPi := new(big.Float).SetPrec(filonRefPrec).Quo(filonPi, fb(2))
if r.Cmp(halfPi) > 0 {
r.Sub(r, filonTwoPi)
} else if r.Cmp(new(big.Float).SetPrec(filonRefPrec).Neg(halfPi)) < 0 {
r.Add(r, filonTwoPi)
}
// Taylor runs about the reduced argument; the zero remainder is the
// exact answer both series converge to.
if r.Sign() == 0 {
return fb(0), fb(1)
}
r2 := new(big.Float).SetPrec(filonRefPrec).Mul(r, r)
ts, tc := new(big.Float).SetPrec(filonRefPrec).Set(r), fb(1)
sumS, sumC := new(big.Float).SetPrec(filonRefPrec).Set(r), fb(1)
for j := int64(1); ; j++ {
ts.Mul(ts, r2)
ts.Quo(ts, fbInt((2*j)*(2*j+1)))
ts.Neg(ts)
sumS.Add(sumS, ts)
tc.Mul(tc, r2)
tc.Quo(tc, fbInt((2*j-1)*(2*j)))
tc.Neg(tc)
sumC.Add(sumC, tc)
if ts.Sign() == 0 || ts.MantExp(nil) < -int(filonRefPrec)-10 {
break
}
}
return sumS, sumC
}
// poly is a real polynomial, coefficients ascending.
type poly []float64
func (p poly) evalBig(x *big.Float) *big.Float {
acc := fb(0)
for _, v := range slices.Backward(p) {
acc.Mul(acc, x)
acc.Add(acc, fb(v))
}
return acc
}
// formalDeriv differentiates the coefficient list.
func (p poly) formalDeriv() poly {
if len(p) <= 1 {
return poly{0}
}
d := make(poly, len(p)-1)
for i := 1; i < len(p); i++ {
d[i-1] = float64(i) * p[i]
}
return d
}
// filonPolyRef evaluates ∫ₐ^b p(x)·cos(kx) dx and the sine part against
// the antiderivative above, in extended precision.
func filonPolyRef(p poly, a, b, k float64) (c, s float64) {
endpoint := func(x float64) (re, im *big.Float) {
// Q(x) = Σ (−1)^j p^{(j)}(x)/(ik)^{j+1}, split by the cycle of
// i^{−(j+1)}: −i, −1, i, 1.
qre, qim := fb(0), fb(0)
sign := fb(1)
kpow := fb(1) // k^(j+1), built by repeated multiplication
px := fb(x)
pp := p
value := pp.evalBig(px)
for j := range p {
kpow.Mul(kpow, fb(k))
scale := new(big.Float).SetPrec(filonRefPrec).Quo(sign, kpow)
switch (j + 1) % 4 {
case 1: // −i
qim.Sub(qim, scale.Mul(scale, value))
case 2: // −1
qre.Sub(qre, scale.Mul(scale, value))
case 3: // i
qim.Add(qim, scale.Mul(scale, value))
default: // 1
qre.Add(qre, scale.Mul(scale, value))
}
sign.Neg(sign)
// p^{(j+1)} for the next term.
pp = pp.formalDeriv()
value = pp.evalBig(px)
}
sb, cb := filonBigSinCos(fb(k * x))
ere := new(big.Float).SetPrec(filonRefPrec).Mul(cb, qre)
ere.Sub(ere, new(big.Float).SetPrec(filonRefPrec).Mul(sb, qim))
eim := new(big.Float).SetPrec(filonRefPrec).Mul(sb, qre)
eim.Add(eim, new(big.Float).SetPrec(filonRefPrec).Mul(cb, qim))
return ere, eim
}
reB, imB := endpoint(b)
reA, imA := endpoint(a)
cv, _ := reB.Sub(reB, reA).Float64()
sv, _ := imB.Sub(imB, imA).Float64()
return cv, sv
}
func filonRelErr(got, want float64) float64 {
if want == 0 {
return math.Abs(got)
}
return math.Abs(got-want) / math.Abs(want)
}
// ampScale bounds |p| over [a, b] by the sum of the coefficients'
// magnitudes lifted to the interval's ends, the scale the absolute
// tolerance is measured against: at high frequency the integral itself
// can cancel to nearly nothing and a relative metric would chase noise.
func ampScale(p poly, a, b float64) float64 {
mag := math.Max(1, math.Max(math.Abs(a), math.Abs(b)))
s := 0.0
pow := 1.0
for _, c := range p {
s += math.Abs(c) * pow
pow *= mag
}
return math.Abs(b-a) * s
}
// TestIntegrateFilonPolynomialMoments holds the scheme against the exact
// antiderivative across amplitudes of every degree the default node
// count interpolates exactly, intervals with both orientations' worth of
// geometry, and frequencies from the settled to the far oscillatory.
func TestIntegrateFilonPolynomialMoments(t *testing.T) {
amplitudes := map[string]poly{
"1": {1},
"2 − 3x + x²": {2, -3, 1},
"1 + 0.5x³": {1, 0, 0, 0.5},
"x − 2x⁴ + 4x⁷": {0, 1, 0, 0, -2, 0, 0, 4},
// Degree 15, the exact interpolation degree of the default 16
// nodes.
"degree 15": {1, -1, 0.5, 0.25, -0.125, 0.0625, 0.5, -0.5, 0.25, -0.25, 0.125, -0.125, 0.0625, -0.0625, 0.5, -0.25},
}
worst := 0.0
for _, k := range []float64{1, 10, 100, 1000, 10000} {
for label, p := range amplitudes {
for _, c := range [][2]float64{{0, 1}, {2, 7}, {-1, 3}} {
a, b := c[0], c[1]
gotC, gotS, err := IntegrateFilon(plainPoly(p), a, b, k, FilonOptions{})
if err != nil {
t.Fatalf("%s on [%g,%g] k=%g: %v", label, a, b, k, err)
}
wantC, wantS := filonPolyRef(p, a, b, k)
scale := ampScale(p, a, b)
dC := math.Abs(gotC-wantC) / scale
dS := math.Abs(gotS-wantS) / scale
worst = math.Max(worst, math.Max(dC, dS))
if dC > 1e-13 {
t.Fatalf("%s on [%g,%g] k=%g: cos part %.17g against exact %.17g (scaled %.3g)",
label, a, b, k, gotC, wantC, dC)
}
if dS > 1e-13 {
t.Fatalf("%s on [%g,%g] k=%g: sin part %.17g against exact %.17g (scaled %.3g)",
label, a, b, k, gotS, wantS, dS)
}
}
}
}
t.Logf("worst scaled moment error across the sweep: %.3e", worst)
}
// plainPoly wraps a polynomial as the amplitude IntegrateFilon samples.
func plainPoly(p poly) func(float64) (float64, error) {
return func(x float64) (float64, error) {
acc := 0.0
for _, v := range slices.Backward(p) {
acc = acc*x + v
}
return acc, nil
}
}
// TestIntegrateFilonLargeFrequency pins the phase at frequencies where a
// float64 antiderivative would stop being a reference: the extended
// precision one keeps counting. The bound is absolute against the
// amplitude scale, because the integral itself shrinks like 1/k.
func TestIntegrateFilonLargeFrequency(t *testing.T) {
amplitudes := map[string]poly{
"2 − 3x + x²": {2, -3, 1},
"1 + 0.5x³": {1, 0, 0, 0.5},
"x − 2x⁴ + 4x⁷": {0, 1, 0, 0, -2, 0, 0, 4},
}
worst := 0.0
for _, k := range []float64{1e5, 1e6} {
for label, p := range amplitudes {
gotC, gotS, err := IntegrateFilon(plainPoly(p), 0, 1, k, FilonOptions{})
if err != nil {
t.Fatalf("%s at k=%g: %v", label, k, err)
}
wantC, wantS := filonPolyRef(p, 0, 1, k)
scale := ampScale(p, 0, 1)
dC, dS := math.Abs(gotC-wantC)/scale, math.Abs(gotS-wantS)/scale
t.Logf("k=%g %s: cos %.3e, sin %.3e (scaled absolute)", k, label, dC, dS)
worst = math.Max(worst, math.Max(dC, dS))
}
}
if worst > 1e-11 {
t.Fatalf("the worst large-frequency scaled error %.3e is past the phase budget", worst)
}
}
// TestIntegrateFilonBeatsPlainQuad measures the scheme's reason to exist:
// at equal evaluation budgets the plain adaptive rule must resolve the
// carrier while Filon tracks the amplitude, and the gap has to be worth
// the second entry point.
func TestIntegrateFilonBeatsPlainQuad(t *testing.T) {
p := poly{0, 1, 0, 0, -2, 0, 0, 4}
for _, k := range []float64{1000, 10000} {
wantC, _ := filonPolyRef(p, 0, 1, k)
// Filon's budget: automatic panels times the default node count.
panels := int(math.Ceil(k / (2 * 16)))
budget := panels * 16
// The plain rule pays 31 evaluations per subinterval (a 21-point
// rule and a 10-point rule on every leaf).
leaves := budget / 31
counts := 0
amp := func(x float64) (float64, error) {
counts++
v, err := plainPoly(p)(x)
if err != nil {
return 0, err
}
return v * math.Cos(k*x), nil
}
got, _, err := IntegrateFunction(amp, 0, 1, QuadratureOptions{MaxIntervals: leaves})
quadErr := math.Inf(1)
if err != nil {
t.Logf("k=%g: the plain rule failed within %d leaves (%d evaluations): %v", k, leaves, counts, err)
} else {
quadErr = filonRelErr(got, wantC)
t.Logf("k=%g: plain quad %.3e (%d evaluations), Filon on the same budget below", k, quadErr, counts)
}
fCounts := 0
f := func(x float64) (float64, error) {
fCounts++
return plainPoly(p)(x)
}
gotC, _, err := IntegrateFilon(f, 0, 1, k, FilonOptions{})
if err != nil {
t.Fatal(err)
}
fErr := filonRelErr(gotC, wantC)
t.Logf("k=%g: Filon %.3e from %d evaluations", k, fErr, fCounts)
if fCounts > budget {
t.Fatalf("Filon spent %d evaluations past its own budget %d", fCounts, budget)
}
if fErr >= quadErr {
t.Fatalf("k=%g: Filon's error %.3e fails to beat the plain rule's %.3e on the same budget", k, fErr, quadErr)
}
if fErr*100 > quadErr {
t.Fatalf("k=%g: Filon's error %.3e is within two orders of the plain rule's %.3e", k, fErr, quadErr)
}
}
}
+68
View File
@@ -0,0 +1,68 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"testing"
)
// TestIntegrateFilonRejectsInfiniteBounds pins the refusal: an infinite
// bound used to fall into an unrepresentable panel count (the conversion
// of the ceiling of infinity) and come back as a quiet NaN pair with a
// nil error, and with k = 0 it sampled the amplitude at infinity. The
// error contract refuses NaN bounds and an infinite frequency; an
// infinite bound is the same breach and answers the same way.
func TestIntegrateFilonRejectsInfiniteBounds(t *testing.T) {
f := func(x float64) (float64, error) { return math.Exp(-x), nil }
cases := []struct {
label string
a, b, k float64
}{
{"upper tail", 0, math.Inf(1), 1},
{"lower tail", math.Inf(-1), 0, 1},
{"whole line", math.Inf(-1), math.Inf(1), 1},
{"zero frequency upper tail", 0, math.Inf(1), 0},
{"span product overflow", 0, 1e100, 1e300},
}
for _, c := range cases {
cos, sin, err := IntegrateFilon(f, c.a, c.b, c.k, FilonOptions{})
if err == nil {
t.Fatalf("%s: [%g, %g] at k=%g integrated to (%g, %g) with no error, want a refusal",
c.label, c.a, c.b, c.k, cos, sin)
}
}
}
// TestIntegrateFilonRejectsUnrepresentablePanels pins the two refusals
// the infinite-bound guard alone does not reach. A panel estimate that
// stays finite but sits beyond the int range converted to garbage: a
// wrapped negative count iterated zero times and answered a quiet zero,
// and a saturated count would iterate forever. And two finite bounds
// far enough apart that their span overflows leave an infinite panel
// width, where a constant amplitude stays finite and the carrier phase
// 0·Inf comes back as a quiet NaN pair. Both answer with the error
// contract instead.
func TestIntegrateFilonRejectsUnrepresentablePanels(t *testing.T) {
constant := func(float64) (float64, error) { return 1, nil }
cases := []struct {
label string
a, b, k float64
opts FilonOptions
f func(x float64) (float64, error)
}{
{"finite estimate beyond the int range", 0, 1e8, 1e300, FilonOptions{}, constant},
{"finite estimate just above the int range", 0, 1e18, 6e2, FilonOptions{}, constant},
{"overflowing span at zero frequency", -math.MaxFloat64, math.MaxFloat64, 0, FilonOptions{}, constant},
{"overflowing span with forced panels", -math.MaxFloat64, math.MaxFloat64, 0, FilonOptions{Panels: 2}, constant},
{"overflowing span with a frequency", -math.MaxFloat64, math.MaxFloat64, 1, FilonOptions{}, constant},
}
for _, c := range cases {
cos, sin, err := IntegrateFilon(c.f, c.a, c.b, c.k, c.opts)
if err == nil {
t.Fatalf("%s: [%g, %g] at k=%g integrated to (%g, %g) with no error, want a refusal",
c.label, c.a, c.b, c.k, cos, sin)
}
}
}
+189
View File
@@ -0,0 +1,189 @@
// 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/base"
)
// intFilon closed forms: the antiderivatives the referents come from.
// intFilon1 integrates 1·cos(kx) and 1·sin(kx); intFilonX integrates x
// against the same kernels; intFilonExp integrates e^{ax} against
// them. All are exact calculus, evaluated independently of the code
// under test.
func intFilon1(a, b, k float64) (c, s float64) {
return (math.Sin(k*b) - math.Sin(k*a)) / k,
(math.Cos(k*a) - math.Cos(k*b)) / k
}
func intFilonX(a, b, k float64) (c, s float64) {
cb, sb := math.Cos(k*b), math.Sin(k*b)
ca, sa := math.Cos(k*a), math.Sin(k*a)
c = (cb+k*b*sb)/k/k - (ca+k*a*sa)/k/k
s = (sb-k*b*cb)/k/k - (sa-k*a*ca)/k/k
return c, s
}
func intFilonExp(a, b, amp, k float64) (c, s float64) {
cb, sb := math.Cos(k*b), math.Sin(k*b)
ca, sa := math.Cos(k*a), math.Sin(k*a)
eb, ea := math.Exp(amp*b), math.Exp(amp*a)
c = eb*(amp*cb+k*sb)/(amp*amp+k*k) - ea*(amp*ca+k*sa)/(amp*amp+k*k)
s = eb*(amp*sb-k*cb)/(amp*amp+k*k) - ea*(amp*sa-k*ca)/(amp*amp+k*k)
return c, s
}
func wantFilon(t *testing.T, label string, gotC, gotS, wantC, wantS, tol float64) {
t.Helper()
if d := math.Abs(gotC - wantC); d > tol*math.Max(1, math.Abs(wantC)) {
t.Fatalf("%s: cos part = %.17g, want %.17g (absolute %.3g)", label, gotC, wantC, d)
}
if d := math.Abs(gotS - wantS); d > tol*math.Max(1, math.Abs(wantS)) {
t.Fatalf("%s: sin part = %.17g, want %.17g (absolute %.3g)", label, gotS, wantS, d)
}
}
// TestIntegrateFilonExactAmplitudes pins the exactness the method
// promises: unit and linear amplitudes are polynomials below the
// default degree, so every frequency from one to a thousand must land
// on the closed form at the rounding floor, whatever the carrier does
// between the samples.
func TestIntegrateFilonExactAmplitudes(t *testing.T) {
one := func(float64) (float64, error) { return 1, nil }
identity := func(x float64) (float64, error) { return x, nil }
for _, c := range []struct {
label string
a, b float64
k float64
f func(float64) (float64, error)
ref func(a, b, k float64) (c, s float64)
}{
{"unit on [0, π]", 0, math.Pi, 1, one, intFilon1},
{"unit on [2, 7]", 2, 7, 100, one, intFilon1},
{"unit on [0, 1]", 0, 1, 1000, one, intFilon1},
{"x on [0, π]", 0, math.Pi, 1, identity, intFilonX},
{"x on [2, 7]", 2, 7, 500, identity, intFilonX},
} {
gotC, gotS, err := IntegrateFilon(c.f, c.a, c.b, c.k, FilonOptions{})
if err != nil {
t.Fatalf("%s: %v", c.label, err)
}
wantC, wantS := c.ref(c.a, c.b, c.k)
wantFilon(t, c.label, gotC, gotS, wantC, wantS, 1e-12)
}
}
// TestIntegrateFilonZeroFrequency pins the degeneration at k = 0: the
// sine part is exactly zero and the cosine part is the plain integral
// of the amplitude.
func TestIntegrateFilonZeroFrequency(t *testing.T) {
f := func(x float64) (float64, error) { return x * x, nil }
gotC, gotS, err := IntegrateFilon(f, 0, 3, 0, FilonOptions{})
if err != nil {
t.Fatalf("IntegrateFilon: %v", err)
}
if gotS != 0 {
t.Fatalf("the sine part at k = 0 is %g, want 0", gotS)
}
if d := math.Abs(gotC - 9); d > 1e-12 {
t.Fatalf("the cosine part at k = 0 is %.17g, want 9", gotC)
}
}
// TestIntegrateFilonExponential holds a non-polynomial amplitude
// against the exact antiderivative at a frequency whose carrier the
// automatic panel count must respect: two hundred and fifty
// oscillations over the interval, answered from a few thousand
// amplitude samples.
func TestIntegrateFilonExponential(t *testing.T) {
const amp = 0.5
f := func(x float64) (float64, error) { return math.Exp(amp * x), nil }
gotC, gotS, err := IntegrateFilon(f, 2, 7, 500, FilonOptions{})
if err != nil {
t.Fatalf("IntegrateFilon: %v", err)
}
wantC, wantS := intFilonExp(2, 7, amp, 500)
wantFilon(t, "exp amplitude", gotC, gotS, wantC, wantS, 1e-11)
}
// TestIntegrateFilonOrientation pins the reversed interval and the
// empty one: reversing negates both parts and an empty interval
// integrates to nothing.
func TestIntegrateFilonOrientation(t *testing.T) {
f := func(x float64) (float64, error) { return math.Exp(0.2 * x), nil }
fc, fs, err := IntegrateFilon(f, 0, 3, 40, FilonOptions{})
if err != nil {
t.Fatal(err)
}
rc, rs, err := IntegrateFilon(f, 3, 0, 40, FilonOptions{})
if err != nil {
t.Fatal(err)
}
if rc != -fc || rs != -fs {
t.Fatalf("the reversed interval gave (%.17g, %.17g), want the negation of (%.17g, %.17g)", rc, rs, fc, fs)
}
if ec, es, err := IntegrateFilon(f, 2, 2, 40, FilonOptions{}); err != nil || ec != 0 || es != 0 {
t.Fatalf("the empty interval gave (%g, %g, %v)", ec, es, err)
}
}
// TestIntegrateFilonDeterministic redraws one integral and requires
// the same bits, the contract every entry point here carries.
func TestIntegrateFilonDeterministic(t *testing.T) {
f := func(x float64) (float64, error) { return math.Exp(0.1 * x), nil }
one := func() (float64, float64) {
c, s, err := IntegrateFilon(f, 0, 5, 300, FilonOptions{})
if err != nil {
t.Fatal(err)
}
return c, s
}
c1, s1 := one()
c2, s2 := one()
if c1 != c2 || s1 != s2 {
t.Fatalf("the same call moved: (%.17g, %.17g) against (%.17g, %.17g)", c1, s1, c2, s2)
}
}
// TestIntegrateFilonErrors pins the contract: NaN bounds or frequency,
// a node count out of range, a forced panel count whose panels carry
// more carrier than the weights can be built within, and a failing or
// non-finite amplitude all surface as errors naming themselves.
func TestIntegrateFilonErrors(t *testing.T) {
f := func(x float64) (float64, error) { return 1, nil }
if _, _, err := IntegrateFilon(f, math.NaN(), 1, 10, FilonOptions{}); err == nil {
t.Fatal("a NaN bound: want an error")
}
if _, _, err := IntegrateFilon(f, 0, 1, math.NaN(), FilonOptions{}); err == nil {
t.Fatal("a NaN frequency: want an error")
}
if _, _, err := IntegrateFilon(f, 0, 1, math.Inf(1), FilonOptions{}); err == nil {
t.Fatal("an infinite frequency: want an error")
}
if _, _, err := IntegrateFilon(f, 0, 1, 10, FilonOptions{Nodes: 1}); err == nil {
t.Fatal("one node: want an error")
}
if _, _, err := IntegrateFilon(f, 0, 1, 10, FilonOptions{Nodes: 33}); err == nil {
t.Fatal("33 nodes: want an error")
}
if _, _, err := IntegrateFilon(f, 0, 1, 1e6, FilonOptions{Panels: 2}); err == nil {
t.Fatal("two panels under a carrier of 1e6: want an error")
}
boom := func(float64) (float64, error) { return 0, base.Errf("amplitude failed") }
if _, _, err := IntegrateFilon(boom, 0, 1, 10, FilonOptions{}); err == nil {
t.Fatal("a failing amplitude: want the error to propagate")
}
bad := func(x float64) (float64, error) {
if x > 0.5 {
return math.NaN(), nil
}
return 1, nil
}
if _, _, err := IntegrateFilon(bad, 0, 1, 10, FilonOptions{}); err == nil {
t.Fatal("a non-finite amplitude value: want an error")
}
}
+165
View File
@@ -0,0 +1,165 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The ADI sweeps enforce the boundary constants on the working state's
// ring rows so the stencils read neighbours unconditionally, and the
// y half-step builds its right sides in column pairs sharing the star
// loads. The published history must not see any of that machinery:
// sample 0 is the initial state exactly, and every later sample matches
// a serial per-line reference bit for bit, non-zero boundary constants
// included (the previous pins all carried zero boundaries, so a swapped
// constant passed silently).
// heat2DReference walks the documented alternating-direction scheme one
// line at a time with fresh scratch: the same expressions in the same
// order the kernel's lanes use, so equal reads give equal bits and the
// comparison below is exact.
func heat2DReference(u0 []float64, rows, cols int, kappa, dx, dy, tFinal, dt float64, samples int, bb, bt, bl, br float64) [][]float64 {
steps, h := pdeSchedule(tFinal, dt, samples)
rx := kappa * h / (2 * dx * dx)
ry := kappa * h / (2 * dy * dy)
u := append([]float64(nil), u0...)
history := [][]float64{append([]float64(nil), u...)}
for c := range cols {
u[c] = bb
u[(rows-1)*cols+c] = bt
}
every := steps / (samples - 1)
star := make([]float64, rows*cols)
lowerX := make([]float64, cols-3)
upperX := make([]float64, cols-3)
diagX := make([]float64, cols-2)
for i := range lowerX {
lowerX[i] = -rx
upperX[i] = -rx
}
for i := range diagX {
diagX[i] = 1 + 2*rx
}
lowerY := make([]float64, rows-3)
upperY := make([]float64, rows-3)
diagY := make([]float64, rows-2)
for i := range lowerY {
lowerY[i] = -ry
upperY[i] = -ry
}
for i := range diagY {
diagY[i] = 1 + 2*ry
}
for s := 1; s <= steps; s++ {
clear(star)
for r := 1; r < rows-1; r++ {
row := u[r*cols : (r+1)*cols]
up := u[(r+1)*cols : (r+2)*cols]
down := u[(r-1)*cols : r*cols]
rhs := make([]float64, cols)
for c := range cols {
rhs[c] = row[c] + ry*(up[c]-2*row[c]+down[c])
}
rhs[1] += rx * bl
rhs[cols-2] += rx * br
dst := star[r*cols+1 : r*cols+cols-1]
err := base.TriSolve(dst, make([]float64, cols-2), make([]float64, cols-2),
lowerX, diagX, upperX, rhs[1:cols-1])
if err != nil {
panic(err)
}
star[r*cols] = bl
star[r*cols+cols-1] = br
}
for c := 1; c < cols-1; c++ {
rhs := make([]float64, rows)
for r := range rows {
off := r * cols
l := star[off+c-1]
wm := star[off+c]
e := star[off+c+1]
rhs[r] = wm + rx*(e-2*wm+l)
}
rhs[1] += ry * bb
rhs[rows-2] += ry * bt
dst := make([]float64, rows-2)
err := base.TriSolve(dst, make([]float64, rows-2), make([]float64, rows-2),
lowerY, diagY, upperY, rhs[1:rows-1])
if err != nil {
panic(err)
}
u[c] = bb
u[(rows-1)*cols+c] = bt
for r := 1; r < rows-1; r++ {
u[r*cols+c] = dst[r-1]
}
}
for r := range rows {
u[r*cols] = bl
u[r*cols+cols-1] = br
}
if s%every == 0 && len(history) < samples {
history = append(history, append([]float64(nil), u...))
}
}
history[samples-1] = append([]float64(nil), u...)
return history
}
func TestHeat2DSampleZeroIsInitialState(t *testing.T) {
// The ring enforcement belongs to the working state: sample 0 is
// the initial state exactly, boundary constants and corners
// included.
const rows, cols = 3, 4
u0 := []float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}
state, err := core.FromFloats(u0, rows, cols)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
hist, err := IntegrateHeat2D(state, 1, 1, 1, 0.1, 0.05, 2, 1, 2, 3, 4)
if err != nil {
t.Fatalf("IntegrateHeat2D: %v", err)
}
got := hist.RawFloats()
for i := range u0 {
if got[i] != u0[i] {
t.Fatalf("sample 0 element %d = %v, want the initial %v", i, got[i], u0[i])
}
}
}
func TestHeat2DSamplesMatchSerialReference(t *testing.T) {
const rows, cols = 5, 7
u0 := make([]float64, rows*cols)
for i := range u0 {
u0[i] = float64(i%13) - 6
}
state, err := core.FromFloats(u0, rows, cols)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
const (
kappa, dx, dy = 0.7, 0.3, 0.25
tFinal, dt = 0.08, 0.01
samples = 4
bb, bt, bl, br = 1.5, -2.25, 3.125, -4.5
)
hist, err := IntegrateHeat2D(state, kappa, dx, dy, tFinal, dt, samples, bb, bt, bl, br)
if err != nil {
t.Fatalf("IntegrateHeat2D: %v", err)
}
want := heat2DReference(u0, rows, cols, kappa, dx, dy, tFinal, dt, samples, bb, bt, bl, br)
got := hist.RawFloats()
for s := range samples {
for i := range u0 {
if got[s*rows*cols+i] != want[s][i] {
t.Fatalf("sample %d element %d = %v, want %v", s, i, got[s*rows*cols+i], want[s][i])
}
}
}
}
+24
View File
@@ -0,0 +1,24 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// mustFloats builds a float array, failing the test on a bad shape.
// Without an explicit shape it defaults to a vector of len(vals).
func mustFloats(t *testing.T, vals []float64, shape ...int) *core.Array {
t.Helper()
if len(shape) == 0 {
shape = []int{len(vals)}
}
a, err := core.FromFloats(vals, shape...)
if err != nil {
t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err)
}
return a
}
+878
View File
@@ -0,0 +1,878 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"errors"
"math"
)
// Ordinary differential equation solvers for initial value problems
// y' = f(t, y). The state y is a rank-1 vector of length n; a system
// of higher rank flattens to its leading-axis vector first.
//
// Three schemes cover the standard regimes. IntegrateODE is the
// workhorse: an adaptive embedded Runge-Kutta pair (Dormand-Prince
// 4(5)) that controls the local error against a mixed absolute and
// relative tolerance. IntegrateRK4 is the classical fixed-step
// fourth-order scheme, useful when a uniform step or simple
// reproducibility per step matters. IntegrateBackwardEuler is the
// entry-level stiff scheme: fully implicit, with each step's
// nonlinear equation solved by Newton over a numerical Jacobian and
// the library's LU solver.
// ODEOptions tunes the adaptive integrator. RelTol ≤ 0 means 1e-6,
// AbsTol ≤ 0 means 1e-9, MaxSteps ≤ 0 means 100000.
type ODEOptions struct {
RelTol float64
AbsTol float64
MaxSteps int
}
// Dormand-Prince 4(5): node offsets, stage coefficients, and the
// 5th- and 4th-order solution weights. Stage 7 shares the 5th-order
// weights (the FSAL property), which is why it needs no separate row.
var (
odeC = [7]float64{0, 1.0 / 5, 3.0 / 10, 4.0 / 5, 8.0 / 9, 1, 1}
odeA = [][]float64{
{},
{1.0 / 5},
{3.0 / 40, 9.0 / 40},
{44.0 / 45, -56.0 / 15, 32.0 / 9},
{19372.0 / 6561, -25360.0 / 2187, 64448.0 / 6561, -212.0 / 729},
{9017.0 / 3168, -355.0 / 33, 46732.0 / 5247, 49.0 / 176, -5103.0 / 18656},
{35.0 / 384, 0, 500.0 / 1113, 125.0 / 192, -2187.0 / 6784, 11.0 / 84},
}
odeB5 = [7]float64{35.0 / 384, 0, 500.0 / 1113, 125.0 / 192, -2187.0 / 6784, 11.0 / 84, 0}
odeB4 = [7]float64{5179.0 / 57600, 0, 7571.0 / 16695, 393.0 / 640, -92097.0 / 339200, 187.0 / 2100, 1.0 / 40}
)
// IntegrateODE integrates y' = f(t, y) from t0 to t1 with the adaptive
// Dormand-Prince 4(5) pair and returns y(t1). Backward integration
// works: a t1 < t0 simply integrates in the negative direction. An
// exhausted step budget, a collapsed step size or an f that returns a
// wrongly shaped state is an error, never a silently truncated
// trajectory.
func IntegrateODE(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, opts ODEOptions) (*core.Array, error) {
return odeRun(f, t0, t1, y0, opts, nil)
}
// readVector copies a's elements into dst, sweeping the raw float64
// payload when a is a dense float64 array and falling back to the
// widening accessor for views and other dtypes. The values written are
// identical either way.
func readVector(dst []float64, a *core.Array) {
if !a.Strided() && a.Dtype() == core.Float {
copy(dst, a.RawFloats())
return
}
for i := range dst {
dst[i] = a.FloatAt(i)
}
}
// denseFloats returns a's elements as a plain float64 slice, sharing
// the payload when a is a dense float64 array and copying the widened
// values otherwise. The values read are the ones the accessor
// returned; a shared slice is read-only, and only an array the caller
// owns may be written through it. A caller sweeping the elements of a
// solver's result uses this instead of one accessor call per element.
func denseFloats(a *core.Array) []float64 {
if !a.Strided() && a.Dtype() == core.Float {
return a.RawFloats()
}
out := make([]float64, a.Len())
for i := range out {
out[i] = a.FloatAt(i)
}
return out
}
// odeRun drives the adaptive Dormand-Prince loop over the whole span.
// When watch is not nil it is called after every accepted step with
// the interval just integrated and clones of the states at both ends;
// a true return stops the integration there, and the watch's error
// aborts it. Everything else behaves exactly like IntegrateODE.
func odeRun(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, opts ODEOptions,
watch func(tPrev, tNow float64, yPrev, yNow []float64) (bool, error)) (*core.Array, error) {
const name = "IntegrateODE"
y, err := odeCheck(name, y0, &opts)
if err != nil {
return nil, err
}
n := len(y)
w := &odeWork{}
w.useStage(n)
k := make([][]float64, 8) // k[1..7] are the stages; k[0] unused
for i := 1; i <= 7; i++ {
k[i] = make([]float64, n)
}
// One scratch accumulator serves every stage of every step: it is
// rebuilt from y at the top of each stage call and read only by
// that call's f evaluation, the same transient view the package's
// fixed-step solvers hand out.
acc := w.stage
stage := func(i int, t float64, h float64) error {
copy(acc, y)
row := odeA[i-1]
for j := 1; j < i; j++ {
if row[j-1] == 0 {
continue
}
// The step-scaled weight is one product, the same (h·a)·k
// grouping the plain component loop evaluated.
hj := h * row[j-1]
kj := k[j]
for m := range n {
acc[m] += hj * kj[m]
}
}
out, err := odeCall(name, f, t+odeC[i-1]*h, acc, n, &w.views)
if err != nil {
return err
}
readVector(k[i], out)
return nil
}
t := t0
h := odeInitialStep(t0, t1, y)
budget := odeBudget{max: opts.MaxSteps}
yEnd := w.yEnd
// The solution weights scaled by the step size: one product each,
// the same (h·b)·k grouping the component loop evaluated.
var hb5, hb4 [7]float64
for !odeArrived(t, t1) {
if err := budget.spend(name, t, t1); err != nil {
return nil, err
}
// Never step past t1; t1−t carries the integration direction.
h = odeClampStep(h, t, t1)
for j := 1; j <= 7; j++ {
hb5[j-1] = h * odeB5[j-1]
hb4[j-1] = h * odeB4[j-1]
}
for i := 1; i <= 7; i++ {
if err := stage(i, t, h); err != nil {
return nil, err
}
}
// The embedded pair: the 5th-order solution advances, the gap
// to the 4th-order one estimates the local error.
errNorm := 0.0
for m := range n {
y5, y4 := y[m], y[m]
for j := 1; j <= 7; j++ {
y5 += hb5[j-1] * k[j][m]
y4 += hb4[j-1] * k[j][m]
}
yEnd[m] = y5
scale := opts.AbsTol + opts.RelTol*math.Max(math.Abs(y[m]), math.Abs(y5))
ratio := (y5 - y4) / scale
errNorm += ratio * ratio
}
errNorm = math.Sqrt(errNorm/float64(n)) + 1e-10
factor := math.Min(5, math.Max(0.2, 0.9*math.Pow(1/errNorm, 1.0/5)))
if errNorm <= 1 {
if watch != nil {
stop, werr := watch(t, t+h, cloneDenseSlice(y), cloneDenseSlice(yEnd))
if werr != nil {
return nil, werr
}
if stop {
return arrayFromVector(yEnd), nil
}
}
copy(y, yEnd)
prevT := t
t += h
h *= factor
// Collapse is "t did not move", not "h is small": a span
// far below the absolute time scale is perfectly
// integrable, and the old absolute floor refused it.
if t == prevT {
return nil, base.Errf("%s: the step size shrank below the resolution of t at t=%g", name, prevT)
}
} else {
// Rejected: retry the same interval with the smaller step.
h *= math.Max(0.2, factor)
}
}
return arrayFromVector(y), nil
}
// IntegrateODEPath integrates y' = f(t, y) from t0 to t1 and returns
// the trajectory sampled at nSamples evenly spaced points, endpoints
// included: times[i] is the sample time and states[i] the state there,
// so states[0] is the initial state and states[nSamples−1] the answer
// IntegrateODE would return. Every interval between neighbouring
// samples is integrated on its own, so the adaptive step control never
// has to align with the sampling grid. Backward integration
// (t1 < t0) works, and the error contract of IntegrateODE applies
// per interval.
func IntegrateODEPath(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, nSamples int, opts ODEOptions) ([]float64, []*core.Array, error) {
if nSamples < 2 {
return nil, nil, base.Errf("IntegrateODEPath: nSamples must be ≥ 2, got %d", nSamples)
}
y, err := odeCheck("IntegrateODEPath", y0, &opts)
if err != nil {
return nil, nil, base.Errf("IntegrateODEPath: %w", err)
}
times := make([]float64, nSamples)
states := make([]*core.Array, nSamples)
times[0] = t0
states[0] = wrapVector(y)
// The last sample is pinned to t1 exactly; the intermediate ones
// are the evenly spaced grid.
for i := 1; i < nSamples; i++ {
times[i] = t0 + float64(i)*(t1-t0)/float64(nSamples-1)
}
times[nSamples-1] = t1
for i := 1; i < nSamples; i++ {
states[i], err = IntegrateODE(f, times[i-1], times[i], states[i-1], opts)
if err != nil {
return nil, nil, base.Errf("IntegrateODEPath: %w", err)
}
}
return times, states, nil
}
// IntegrateODESteps integrates y' = f(t, y) from t0 to t1 and returns
// the trajectory as recorded at every accepted solver step: times[i]
// carries states[i] = y(times[i]), starting with (t0, y0) and ending
// with (t1, y(t1)). The accepted steps are where the adaptive control
// judged the local error within tolerance, so they are the natural
// interpolation nodes for post-processing, sensitivity analysis and
// adjoint passes. Backward integration records descending times; the
// error contract of IntegrateODE applies.
func IntegrateODESteps(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, opts ODEOptions) ([]float64, []*core.Array, error) {
y, err := odeCheck("IntegrateODESteps", y0, &opts)
if err != nil {
return nil, nil, base.Errf("IntegrateODESteps: %w", err)
}
times := []float64{t0}
states := []*core.Array{wrapVector(y)}
watch := func(tPrev, tNow float64, yPrev, yNow []float64) (bool, error) {
times = append(times, tNow)
// yNow is the run's own per-call clone; nothing aliases it.
states = append(states, wrapVector(yNow))
return false, nil
}
if _, err := odeRun(f, t0, t1, y0, opts, watch); err != nil {
return nil, nil, base.Errf("IntegrateODESteps: %w", err)
}
// The run's last accepted boundary is t+h with h = t1−t, which
// rounds a few ulps off t1 whenever the magnitudes demand it; the
// documented endpoint is t1 exactly, and the recorded state there
// is already the run's answer y(t1).
times[len(times)-1] = t1
return times, states, nil
}
// IntegrateRK4 integrates y' = f(t, y) with the classical fixed-step
// fourth-order Runge-Kutta scheme over the given number of equal
// steps, returning y(t1).
func IntegrateRK4(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, steps int) (*core.Array, error) {
if steps <= 0 {
return nil, base.Errf("IntegrateRK4: steps must be ≥ 1, got %d", steps)
}
y, err := odeCheck("IntegrateRK4", y0, nil)
if err != nil {
return nil, err
}
n := len(y)
h := (t1 - t0) / float64(steps)
k1 := make([]float64, n)
k2 := make([]float64, n)
k3 := make([]float64, n)
k4 := make([]float64, n)
tmp := make([]float64, n)
views := &odeViews{}
call := func(t float64, v []float64, out []float64) error {
o, err := odeCall("IntegrateRK4", f, t, v, n, views)
if err != nil {
return err
}
readVector(out, o)
// A non-finite stage flows straight into the state with no
// rejection mechanism to catch it, and the fixed-step run
// would publish NaN with a nil error; the adaptive drivers
// reject it, this one has to refuse it.
for i := range n {
if math.IsNaN(out[i]) || math.IsInf(out[i], 0) {
return base.Errf("IntegrateRK4: f returned the non-finite value %g at coordinate %d, t=%g", out[i], i, t)
}
}
return nil
}
// The stage times come from the exact grid t0 + i·h, never from an
// accumulated t += h: the addition's rounding walks over a long run
// (measured on y' = cos t from t0 = 1e6 the walk contributes an
// error of 1.1e-8 that no step count refines away), while each grid
// point carries a single rounding that stays put.
for i := range steps {
t := t0 + float64(i)*h
if err := call(t, y, k1); err != nil {
return nil, err
}
for i := range n {
tmp[i] = y[i] + h*k1[i]/2
}
if err := call(t+h/2, tmp, k2); err != nil {
return nil, err
}
for i := range n {
tmp[i] = y[i] + h*k2[i]/2
}
if err := call(t+h/2, tmp, k3); err != nil {
return nil, err
}
for i := range n {
tmp[i] = y[i] + h*k3[i]
}
if err := call(t+h, tmp, k4); err != nil {
return nil, err
}
for i := range n {
y[i] += h * (k1[i] + 2*k2[i] + 2*k3[i] + k4[i]) / 6
}
}
return arrayFromVector(y), nil
}
// IntegrateBackwardEuler integrates y' = f(t, y) with the fully
// implicit Euler scheme y_{n+1} = y_n + h·f(t_{n+1}, y_{n+1}), solving
// each step by Newton over a numerical Jacobian and the library's LU
// solver. The extra work per step is what buys stability on stiff
// systems, where the explicit schemes need step sizes far below what
// accuracy alone would ask for.
func IntegrateBackwardEuler(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, steps int, opts ODEOptions) (*core.Array, error) {
if steps <= 0 {
return nil, base.Errf("IntegrateBackwardEuler: steps must be ≥ 1, got %d", steps)
}
y, err := odeCheck("IntegrateBackwardEuler", y0, &opts)
if err != nil {
return nil, err
}
n := len(y)
h := (t1 - t0) / float64(steps)
yn := cloneDenseSlice(y)
seed := make([]float64, n)
fy := make([]float64, n)
// One Newton result buffer serves every step: it aliases neither the
// state nor the seed, and each step overwrites it fully.
zbuf := make([]float64, n)
w := &odeWork{}
// The step times come from the exact grid t0 + i·h, never from an
// accumulated t += h: the addition's rounding walks over a long
// run, while each grid point carries a single rounding that stays
// put.
for i := range steps {
tNext := t0 + float64(i+1)*h
// Newton on G(z) = z − y_n − h·f(t_{n+1}, z) = 0, seeded with
// the semi-implicit Euler prediction.
out, ferr := odeCall("IntegrateBackwardEuler", f, tNext, yn, n, &w.views)
if ferr != nil {
return nil, ferr
}
readVector(fy, out)
for i := range n {
seed[i] = yn[i] + h*fy[i]
}
if nerr := odeNewton("IntegrateBackwardEuler", f, w, tNext, 1, h, yn, seed,
zbuf, opts.AbsTol, opts.RelTol); nerr != nil {
return nil, nerr
}
copy(yn, zbuf)
}
return arrayFromVector(yn), nil
}
// errNewtonStalled marks an implicit solve whose Newton iteration ran
// out of budget or hit a singular matrix without converging. A driver
// that can shrink the step retries on it; a failed f evaluation is a
// different, fatal error.
var errNewtonStalled = errors.New("the Newton iteration did not converge")
// odeCall evaluates f at (t, v) and validates that the result is a
// vector of the expected length n, returning it unchanged. views,
// when not nil, caches the read-only wrapper handed to f, which a
// driver calling f repeatedly wants; a cold call site passes nil.
// Callers that need plain floats follow up with odeEval; callers that
// want to place the values themselves read them straight off the
// array.
func odeCall(name string, f func(t float64, y *core.Array) (*core.Array, error),
t float64, v []float64, n int, views *odeViews) (*core.Array, error) {
out, err := f(t, views.of(v))
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
if out.NDim() != 1 || out.Len() != n {
return nil, base.Errf("%s: f returned shape %s, want a vector of length %d",
name, base.ShapeText(out.Shape()), n)
}
return out, nil
}
// odeEval calls f at (t, v) and returns the derivative as a plain
// float64 slice.
func odeEval(name string, f func(t float64, y *core.Array) (*core.Array, error),
t float64, v []float64, n int, views *odeViews) ([]float64, error) {
out, err := odeCall(name, f, t, v, n, views)
if err != nil {
return nil, err
}
r := make([]float64, n)
readVector(r, out)
return r, nil
}
// odeBudget counts attempted solver steps against the MaxSteps option:
// the budget is spent before the step's stages are evaluated, so a
// rejected step consumes it like an accepted one.
type odeBudget struct {
used int
max int
}
// spend spends one step of the budget, failing once it is exhausted.
func (b *odeBudget) spend(name string, t, t1 float64) error {
if b.used >= b.max {
return base.Errf("%s: reached MaxSteps=%d at t=%g before t1=%g", name, b.max, t, t1)
}
b.used++
return nil
}
// odeClampStep caps h so a step never overshoots t1; t1−t carries the
// integration direction.
func odeClampStep(h, t, t1 float64) float64 {
if math.Abs(h) > math.Abs(t1-t) {
return t1 - t
}
return h
}
// odeArrived reports whether t sits within a few ulps of t1. The
// accumulating t += h can miss the exact endpoint by rounding once
// t and t1 differ in magnitude beyond Sterbenz territory, and the
// residual distance is indistinguishable from zero at working
// precision, so the solvers treat it as arrived rather than report a
// collapsed step over it.
func odeArrived(t, t1 float64) bool {
if t == t1 {
return true
}
return math.Abs(t1-t) <= 8*base.EpsF*math.Max(math.Abs(t), math.Abs(t1))
}
// odeWork holds the scratch the implicit schemes reuse across the
// steps of one run: the Newton vectors, one flat numerical Jacobian
// with its difference stencils, the factored matrix's rows, the ROS4
// stage buffers and the read-only views handed to f. Every buffer is
// overwritten before it is read, so a run builds the workspace once
// and no step allocates scratch of its own.
type odeWork struct {
// The Newton iteration's residual, step and derivative buffers,
// plus the DAE mass-matrix product M·z.
g, col, fzs []float64
mz []float64
// jac is the numerical Jacobian as one flat n×n buffer, row-major:
// jac[i*n+j] is ∂f_i/∂z_j.
jac []float64
// The Jacobian's central-difference stencils: the two perturbed
// states and their two results.
zp, zm, fp, fm []float64
// mat is the implicit relation's matrix in row-major rows, ROS4's
// (1/(γh))I − J or Newton's α·I − h·J depending on the caller, with
// perm the row permutation its factorisation produced. Factor
// shuffles the rows in place, and every rebuild rewrites the lot.
mat [][]float64
perm []int
// The ROS4 stage buffers: the four divided differences, the stage
// value, its right side, f's result and the candidate end state.
ks [4][]float64
stage, rhs, fy []float64
yEnd []float64
// The Newton iterate the implicit schemes iterate in, and the visit
// bitmap the stage permutation walks: both reused across every step
// and every attempt of one run.
zwork []float64
visited []bool
// views caches the wrapper handed to f per scratch slice.
views odeViews
}
// use returns the workspace's buffers at the state length n, growing
// them on first use. The content is left as the previous step wrote
// it: every consumer overwrites its buffer before reading it.
func (w *odeWork) use(n int) {
w.g = sizedBuf(w.g, n)
w.col = sizedBuf(w.col, n)
w.fzs = sizedBuf(w.fzs, n)
w.mz = sizedBuf(w.mz, n)
w.jac = sizedBuf(w.jac, n*n)
w.zp = sizedBuf(w.zp, n)
w.zm = sizedBuf(w.zm, n)
w.fp = sizedBuf(w.fp, n)
w.fm = sizedBuf(w.fm, n)
w.mat = sizedRows(w.mat, n)
for s := range w.ks {
w.ks[s] = sizedBuf(w.ks[s], n)
}
w.stage = sizedBuf(w.stage, n)
w.rhs = sizedBuf(w.rhs, n)
w.fy = sizedBuf(w.fy, n)
w.yEnd = sizedBuf(w.yEnd, n)
w.zwork = sizedBuf(w.zwork, n)
w.visited = sizedBools(w.visited, n)
}
// useStage returns the explicit step loop's buffers at the state
// length n, growing them on first use: the stage accumulator and the
// candidate end state are all an explicit pair needs, and sizing the
// implicit buffers here would allocate the Jacobian and the factored
// matrix for a loop that never takes a derivative.
func (w *odeWork) useStage(n int) {
w.stage = sizedBuf(w.stage, n)
w.yEnd = sizedBuf(w.yEnd, n)
}
// sizedBuf returns b cut to length n, reusing its storage when it is
// large enough.
func sizedBuf(b []float64, n int) []float64 {
if cap(b) < n {
return make([]float64, n)
}
return b[:n]
}
// sizedBools returns b cut to length n, reusing its storage when it is
// large enough.
func sizedBools(b []bool, n int) []bool {
if cap(b) < n {
return make([]bool, n)
}
return b[:n]
}
// odePermuteColumn reorders col in place so that col[i] takes the
// value that sat at perm[i], the permutation the workspace's LU
// factorisation produced. The walk carries each displaced value around
// its cycle exactly as the library's PermuteColumn does, plain
// assignments moving each value once, so the column ends up bit-
// identical; the difference is that the visit bitmap is the caller's
// reused scratch rather than a fresh allocation per call. The bitmap
// is cleared on entry, so a dirty buffer behaves exactly like a fresh
// one.
func odePermuteColumn(col []float64, perm []int, visited []bool) {
vis := visited[:len(col)]
clear(vis)
for i := range col {
if vis[i] || perm[i] == i {
vis[i] = true
continue
}
// Carry the displaced value around the cycle.
tmp := col[i]
j := i
for {
vis[j] = true
k := perm[j]
if k == i {
break
}
col[j] = col[k]
j = k
}
col[j] = tmp
}
}
// arrayFromVector copies a float64 slice into a fresh rank-1 float64
// array: the trajectory endpoint contract, without the intermediate
// wrapper a cloneArray(wrapVector(...)) pair built. The result never
// aliases the input.
func arrayFromVector(v []float64) *core.Array {
a := core.New(core.Float, len(v))
copy(a.RawFloats(), v)
return a
}
// sizedRows returns m as n row slices of length n, reusing the rows it
// already holds. The rows' content is the caller's to overwrite.
func sizedRows(m [][]float64, n int) [][]float64 {
if cap(m) < n {
m = make([][]float64, n)
}
m = m[:n]
for i := range m {
m[i] = sizedBuf(m[i], n)
}
return m
}
// odeViews caches the read-only wrapper handed to f for one scratch
// slice, so a driver that calls f thousands of times builds the
// wrapper once per slice instead of once per call. The values behind
// the wrapper are the driver's own scratch and keep changing exactly
// as they did; only the Array header is reused. The slices a driver
// hands in are few, so a linear scan beats a map; a driver that
// presents a fresh slice every call cannot grow the cache without
// bound, because the oldest entry makes way.
type odeViews struct {
entries []odeView
}
type odeView struct {
vals []float64
arr *core.Array
}
// odeViewSlots bounds the cache. A driver holds a handful of scratch
// slices at once, and each slot pins one sliced buffer, so the bound
// keeps both the scan and the retention small.
const odeViewSlots = 8
// of returns a read-only view of s, reusing the one already built for
// that slice. A nil cache builds a fresh view, which is what a cold
// call site wants.
func (v *odeViews) of(s []float64) *core.Array {
if v == nil || len(s) == 0 {
return wrapVector(s)
}
for i := range v.entries {
e := &v.entries[i]
if len(e.vals) == len(s) && &e.vals[0] == &s[0] {
return e.arr
}
}
arr := wrapVector(s)
if len(v.entries) == odeViewSlots {
copy(v.entries, v.entries[1:])
v.entries = v.entries[:odeViewSlots-1]
}
v.entries = append(v.entries, odeView{vals: s, arr: arr})
return arr
}
// odeJacobian fills the workspace's flat Jacobian with the central
// differences of f at (t, z), one column per state component:
// entry i*n+j is ∂f_i/∂z_j. The returned slice is the workspace's, so
// it stays valid until the next Jacobian. The two perturbed stencils
// and their two result buffers are reused across columns: each round
// rebuilds the stencils from z and overwrites both results before
// reading them.
func odeJacobian(name string, f func(t float64, y *core.Array) (*core.Array, error),
t float64, z []float64, w *odeWork) ([]float64, error) {
n := len(z)
w.use(n)
jac := w.jac
for j := range n {
eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(z[j]))
copy(w.zp, z)
copy(w.zm, z)
w.zp[j] += eps
w.zm[j] -= eps
out, e1 := odeCall(name, f, t, w.zp, n, &w.views)
if e1 != nil {
return nil, e1
}
readVector(w.fp, out)
out, e2 := odeCall(name, f, t, w.zm, n, &w.views)
if e2 != nil {
return nil, e2
}
readVector(w.fm, out)
for i := range n {
jac[i*n+j] = (w.fp[i] - w.fm[i]) / (2 * eps)
}
}
return jac, nil
}
// odeNewton solves the implicit step equation α·z − h·f(tNext, z) = β
// for z by Newton over a numerical Jacobian and the library's LU
// solver, writing the converged state into dst and returning an error
// otherwise. The Jacobian is frozen from the seed and rebuilt twice
// when convergence drags, so a converged step costs one Jacobian and a
// handful of f evaluations. The iteration runs in the workspace's own
// buffer and the column permutation walks the workspace's bitmap, so
// the only allocation a converged step costs is the caller's dst: a
// driver whose dst recycles through a ring allocates its result
// buffers once per solve, not once per step, and a rejected or stalled
// attempt allocates nothing. dst must not alias seed or beta; the
// workspace overwrites it fully at convergence. Convergence is
// measured on the residual against the error scale the caller
// integrates to, an order of magnitude below it, but never below the
// floating-point floor of the residual's own terms, which would
// otherwise be unreachable at the tiny steps a stiff start begins
// with. An iteration that outlives twenty rounds, or a singular Newton
// matrix, surfaces as errNewtonStalled so a stepping driver can retry
// with a smaller step; an f that fails is the fatal error it is.
func odeNewton(name string, f func(t float64, y *core.Array) (*core.Array, error),
w *odeWork, tNext float64, alpha, h float64, beta, seed, dst []float64,
absTol, relTol float64) error {
n := len(seed)
w.use(n)
// The iterate starts as the seed and stays the workspace's buffer:
// the views cache then hands f one stable wrapper for every Newton
// call of the run.
z := sizedBuf(w.zwork, n)
copy(z, seed)
for iteration := range 20 {
out, err := odeCall(name, f, tNext, z, n, &w.views)
if err != nil {
return err
}
readVector(w.fzs, out)
worst, terms := 0.0, 0.0
for i := range n {
w.g[i] = alpha*z[i] - h*w.fzs[i] - beta[i]
worst = math.Max(worst, math.Abs(w.g[i]))
terms = math.Max(terms, math.Abs(alpha*z[i])+math.Abs(h*w.fzs[i])+math.Abs(beta[i]))
}
limit := math.Max(0.1*(absTol+relTol*normInfOfStep(z)), 8*base.EpsF*terms)
if worst <= limit {
copy(dst, z)
return nil
}
if iteration == 0 || iteration == 4 || iteration == 10 {
jac, jerr := odeJacobian(name, f, tNext, z, w)
if jerr != nil {
return jerr
}
// Newton matrix α·I − h·J, a fresh LU for the frozen
// Jacobian; the iterations that follow only substitute.
// Every row is rebuilt entry by entry before the
// factorisation reads it.
for i := range n {
row := w.mat[i]
for j := range n {
row[j] = -h * jac[i*n+j]
}
row[i] += alpha
}
w.perm, _ = base.Factor(w.mat)
if err := base.CheckSingular(name, w.mat); err != nil {
return base.Errf("%s: %w, singular Newton matrix at t=%g",
name, errNewtonStalled, tNext)
}
}
for i := range n {
w.col[i] = -w.g[i]
}
odePermuteColumn(w.col, w.perm, w.visited)
base.SolveColumn(w.mat, w.col)
for i := range n {
z[i] += w.col[i]
}
}
return base.Errf("%s: %w at t=%g", name, errNewtonStalled, tNext)
}
// odeCheck validates the initial state, applies option defaults and
// returns the flat float64 working state.
func odeCheck(name string, y0 *core.Array, opts *ODEOptions) ([]float64, error) {
if y0.NDim() != 1 {
return nil, base.Errf("%s: the state must be a vector, got shape %s", name, base.ShapeText(y0.Shape()))
}
if y0.Len() == 0 {
return nil, base.Errf("%s: the state must not be empty", name)
}
if y0.Dtype() == core.Complex {
return nil, base.Errf("%s: complex states are not supported", name)
}
if opts != nil {
if opts.RelTol <= 0 {
opts.RelTol = 1e-6
}
if opts.AbsTol <= 0 {
opts.AbsTol = 1e-9
}
if opts.MaxSteps <= 0 {
opts.MaxSteps = 100000
}
}
return cloneDense(y0), nil
}
// odeInitialStep guesses the first step size as a small fraction of
// the integration span, carrying the direction in its sign.
func odeInitialStep(t0, t1 float64, y []float64) float64 {
h := 0.01 * math.Abs(t1-t0)
if h == 0 {
h = 1e-6
}
// The step sign carries the integration direction: a t1 < t0
// integrates backwards.
if t1 < t0 {
h = -h
}
return h
}
// cloneArray copies an array element by element, so the integration
// steps' results never alias a buffer already handed out.
func cloneArray(a *core.Array) *core.Array {
out := core.New(a.Dtype(), a.Shape()...)
switch a.Dtype() {
case core.Float:
copy(out.RawFloats(), a.RawFloats())
case core.Float32:
copy(out.RawFloat32s(), a.RawFloat32s())
case core.Int:
copy(out.RawInts(), a.RawInts())
default:
copy(out.RawComplexes(), a.RawComplexes())
}
return out
}
// normInfOfStep returns the infinity norm of a step vector.
func normInfOfStep(step []float64) float64 {
worst := 0.0
for _, v := range step {
if a := math.Abs(v); a > worst {
worst = a
}
}
return worst
}
// wrapVector views a float64 slice as a rank-1 Array without copying.
// The caller must treat the result as read-only.
func wrapVector(v []float64) *core.Array {
a, _ := core.FloatsFromArray(v, len(v))
return a
}
// cloneDense copies an array's elements into a plain float64 slice,
// widening int and float32 elements exactly.
func cloneDense(y0 *core.Array) []float64 {
vals := make([]float64, y0.Len())
for i := range vals {
vals[i] = y0.FloatAt(i)
}
return vals
}
// cloneDenseSlice copies a float64 slice.
func cloneDenseSlice(v []float64) []float64 {
out := make([]float64, len(v))
copy(out, v)
return out
}
+475
View File
@@ -0,0 +1,475 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"math"
"testing"
)
// decay returns f for y' = −y, the reference every scheme must nail.
func decay(t float64, y *core.Array) (*core.Array, error) {
return core.MulF(y, -1), nil
}
// TestIntegrateODEExponential checks the adaptive solver against the
// analytic exponential decay, forward and backward in time.
func TestIntegrateODEExponential(t *testing.T) {
y0 := mustFloats(t, []float64{1})
end, err := IntegrateODE(decay, 0, 1, y0, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODE: %v", err)
}
if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-9 {
t.Fatalf("y(1) = %.14g, want %.14g", end.FloatAt(0), math.Exp(-1))
}
// Backward integration from t=1 to t=0 must return the start.
start, err := IntegrateODE(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}),
ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODE backward: %v", err)
}
if math.Abs(start.FloatAt(0)-1) > 1e-8 {
t.Fatalf("backward y(0) = %.14g, want 1", start.FloatAt(0))
}
}
// TestIntegrateODEOscillator checks a two-dimensional linear system
// against the analytic phase rotation.
func TestIntegrateODEOscillator(t *testing.T) {
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
y0 := mustFloats(t, []float64{1, 0})
end, err := IntegrateODE(f, 0, math.Pi/2, y0, ODEOptions{RelTol: 1e-11, AbsTol: 1e-13})
if err != nil {
t.Fatalf("IntegrateODE: %v", err)
}
// y1 = cos t, y2 = −sin t, so a quarter period lands on (0, −1).
if math.Abs(end.FloatAt(0)) > 1e-7 || math.Abs(end.FloatAt(1)+1) > 1e-7 {
t.Fatalf("quarter period = (%.10g, %.10g), want (0, -1)",
end.FloatAt(0), end.FloatAt(1))
}
// A full period returns to the start.
full, err := IntegrateODE(f, 0, 2*math.Pi, y0, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODE full: %v", err)
}
for i := range 2 {
if math.Abs(full.FloatAt(i)-y0.FloatAt(i)) > 1e-6 {
t.Fatalf("full period[%d] = %v, want %v", i, full.FloatAt(i), y0.FloatAt(i))
}
}
}
// TestIntegrateRk4 checks the fixed-step scheme's fourth-order
// convergence: halving h must shrink the error by roughly sixteen.
func TestIntegrateRk4(t *testing.T) {
errAt := func(steps int) float64 {
end, err := IntegrateRK4(decay, 0, 1, mustFloats(t, []float64{1}), steps)
if err != nil {
t.Fatalf("IntegrateRK4(%d): %v", steps, err)
}
return math.Abs(end.FloatAt(0) - math.Exp(-1))
}
e10, e20 := errAt(10), errAt(20)
if e10 < 1e-13 {
t.Skipf("error already at round-off (%v)", e10)
}
ratio := e10 / e20
if ratio < 12 || ratio > 20 {
t.Fatalf("error ratio over a halved step = %.2g, want ≈ 16 for a fourth-order scheme", ratio)
}
}
// TestIntegrateBackwardEulerStiff demonstrates the reason an implicit
// scheme exists: on y' = −1000y the implicit Euler stays bounded and
// matches its closed-form damping at a step where explicit schemes
// blow up.
func TestIntegrateBackwardEulerStiff(t *testing.T) {
const lambda = 1000.0
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.MulF(y, -lambda), nil
}
// Ten steps of h = 0.01: h·λ = 10, far outside RK4's stability
// region but perfectly damped for the implicit scheme.
y0 := mustFloats(t, []float64{1})
end, err := IntegrateBackwardEuler(f, 0, 0.1, y0, 10, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateBackwardEuler: %v", err)
}
// Closed form of one implicit Euler step: y_{n+1} = y_n/(1+hλ),
// so y_10 = (1/11)^10.
damped := math.Pow(1/(1+0.01*lambda), 10)
if math.Abs(end.FloatAt(0)-damped) > 1e-9*math.Abs(damped) {
t.Fatalf("stiff result = %.12g, want %.12g", end.FloatAt(0), damped)
}
if end.FloatAt(0) <= 0 {
t.Fatalf("the implicit scheme must stay positive on decay, got %v", end.FloatAt(0))
}
}
// TestIntegrateBackwardEulerAccuracy checks that with a sane step the
// implicit scheme tracks the analytic decay as well.
func TestIntegrateBackwardEulerAccuracy(t *testing.T) {
end, err := IntegrateBackwardEuler(decay, 0, 1, mustFloats(t, []float64{1}), 1000, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateBackwardEuler: %v", err)
}
// First order: the global error is O(h), about h/2 for decay.
if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-3 {
t.Fatalf("y(1) = %.14g, want %.14g ± 1e-3", end.FloatAt(0), math.Exp(-1))
}
}
// TestIntegrateBackwardEulerVectorState pins the Newton solve on a
// two-dimensional stiff system: both modes decay at their own rate,
// and constant steps give each component the closed form
// y(t1) = y0·(1+h·λ)^{−steps} an implicit Euler step has on y' = −λy.
func TestIntegrateBackwardEulerVectorState(t *testing.T) {
const slow = 1.0
const fast = 1000.0
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{-fast * y.FloatAt(0), -slow * y.FloatAt(1)}, 2)
}
y0 := mustFloats(t, []float64{1, 1})
end, err := IntegrateBackwardEuler(f, 0, 1, y0, 10, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateBackwardEuler: %v", err)
}
wantFast := math.Pow(1/(1+0.1*fast), 10)
if math.Abs(end.FloatAt(0)-wantFast) > 1e-9*math.Abs(wantFast) {
t.Fatalf("fast mode = %.14g, want %.14g", end.FloatAt(0), wantFast)
}
wantSlow := math.Pow(1/(1+0.1*slow), 10)
if math.Abs(end.FloatAt(1)-wantSlow) > 1e-9*math.Abs(wantSlow) {
t.Fatalf("slow mode = %.14g, want %.14g", end.FloatAt(1), wantSlow)
}
}
// TestODEErrors pins the error contracts shared by the solvers.
func TestODEErrors(t *testing.T) {
wrongShape := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{1, 1}, 2)
}
y0 := mustFloats(t, []float64{1})
if _, err := IntegrateODE(wrongShape, 0, 1, y0, ODEOptions{}); err == nil {
t.Fatal("expected an error when f returns the wrong shape")
}
if _, err := IntegrateODE(wrongShape, 0, 1, y0, ODEOptions{}); err == nil {
t.Fatal("expected an error when f returns the wrong shape in RK4 as well")
}
if _, err := IntegrateRK4(decay, 0, 1, y0, 0); err == nil {
t.Fatal("expected an error for zero steps")
}
matrixState := mustFloats(t, []float64{1, 1}, 1, 2)
if _, err := IntegrateODE(decay, 0, 1, matrixState, ODEOptions{}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
empty := mustFloats(t, nil)
if _, err := IntegrateODE(decay, 0, 1, empty, ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
// A tight budget on a slow decay must report the budget, not lie.
if _, err := IntegrateODE(decay, 0, 1, y0, ODEOptions{MaxSteps: 3}); err == nil {
t.Fatal("expected an error for an exhausted step budget")
}
}
// TestIntegrateODEPathExponential checks the sampled trajectory of
// y' = −y against the analytic decay at every sample point.
func TestIntegrateODEPathExponential(t *testing.T) {
const n = 5
times, states, err := IntegrateODEPath(decay, 0, 2, mustFloats(t, []float64{1}),
n, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODEPath: %v", err)
}
if len(times) != n || len(states) != n {
t.Fatalf("lengths (%d, %d), want (%d, %d)", len(times), len(states), n, n)
}
for i := range n {
want := 0.5 * float64(i)
if math.Abs(times[i]-want) > 1e-12 {
t.Fatalf("times[%d] = %.14g, want %.14g", i, times[i], want)
}
got := states[i].FloatAt(0)
if math.Abs(got-math.Exp(-want)) > 1e-9 {
t.Fatalf("y(%.1f) = %.14g, want %.14g", want, got, math.Exp(-want))
}
}
}
// TestIntegrateODEPathBackward samples a backwards integration: the
// times descend and the analytic law still holds per sample.
func TestIntegrateODEPathBackward(t *testing.T) {
times, states, err := IntegrateODEPath(decay, 2, 0, mustFloats(t, []float64{math.Exp(-2)}),
3, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODEPath backward: %v", err)
}
for i, want := range []float64{2, 1, 0} {
if math.Abs(times[i]-want) > 1e-12 {
t.Fatalf("times[%d] = %.14g, want %g", i, times[i], want)
}
if math.Abs(states[i].FloatAt(0)-math.Exp(-want)) > 1e-9 {
t.Fatalf("y(%g) = %.14g, want %.14g", want, states[i].FloatAt(0), math.Exp(-want))
}
}
}
// TestIntegrateODEPathStatesIndependent checks the returned states do
// not alias one another or the initial vector: a later integration
// must never rewrite an earlier sample.
func TestIntegrateODEPathStatesIndependent(t *testing.T) {
y0 := mustFloats(t, []float64{1, 0})
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
_, states, err := IntegrateODEPath(f, 0, 1, y0, 4, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODEPath: %v", err)
}
states[3].RawFloats()[0] = 99 // must not leak into y0 or the other samples
if y0.FloatAt(0) != 1 {
t.Fatal("mutating a sample changed the caller's initial state")
}
if states[2].FloatAt(0) == 99 {
t.Fatal("samples alias one another")
}
}
// TestIntegrateODEPathErrors pins the path-specific error contract.
func TestIntegrateODEPathErrors(t *testing.T) {
y0 := mustFloats(t, []float64{1})
if _, _, err := IntegrateODEPath(decay, 0, 1, y0, 1, ODEOptions{}); err == nil {
t.Fatal("expected an error for a single sample")
}
// An f that fails past the midpoint must fail the whole path.
boom := func(t float64, y *core.Array) (*core.Array, error) {
if t > 0.6 {
return nil, base.Errf("detector tripped")
}
return core.MulF(y, -1), nil
}
if _, _, err := IntegrateODEPath(boom, 0, 1, y0, 5, ODEOptions{}); err == nil {
t.Fatal("expected the operator error to propagate")
}
if _, _, err := IntegrateODEPath(decay, 0, 1, mustFloats(t, nil), 3, ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
}
// TestIntegrateODEEventsProjectile drops a projectile and watches its
// height: the crossing of zero is the flight time 2v₀/g, analytically
// known, with the state's velocity the exact mirror of the launch.
func TestIntegrateODEEventsProjectile(t *testing.T) {
const g = 9.81
const v0 = 10
f := func(now float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{y.FloatAt(1), -g}), nil
}
height := func(now float64, y *core.Array) (float64, error) {
return y.FloatAt(0), nil
}
hits, final, err := IntegrateODEEvents(f, 0, 5, mustFloats(t, []float64{0, v0}),
[]ODEWatch{{Function: height, Direction: -1}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents: %v", err)
}
if len(hits) != 1 {
t.Fatalf("hits = %d, want 1", len(hits))
}
wantT := 2 * v0 / g
if math.Abs(hits[0].Time-wantT) > 1e-9 {
t.Fatalf("impact at %.14g, want %.14g", hits[0].Time, wantT)
}
if math.Abs(hits[0].State.FloatAt(1)+v0) > 1e-7 {
t.Fatalf("impact speed %.10g, want %.10g", hits[0].State.FloatAt(1), float64(-v0))
}
if hits[0].Rising {
t.Fatal("height crossing on the way down must not be rising")
}
if final.Len() != 2 {
t.Fatalf("final state shape %v", final.Shape())
}
}
// TestIntegrateODEEventsOscillator watches the oscillator's position
// over five periods: cos crosses zero once per half period, and the
// falling-only filter keeps every second one, at π/2 + 2πk.
func TestIntegrateODEEventsOscillator(t *testing.T) {
f := func(now float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{y.FloatAt(1), -y.FloatAt(0)}), nil
}
position := func(now float64, y *core.Array) (float64, error) {
return y.FloatAt(0), nil
}
hits, _, err := IntegrateODEEvents(f, 0, 5*2*math.Pi, mustFloats(t, []float64{1, 0}),
[]ODEWatch{{Function: position, Direction: -1}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents: %v", err)
}
if len(hits) != 5 {
t.Fatalf("hits = %d, want 5 falling crossings", len(hits))
}
for k, hit := range hits {
want := math.Pi/2 + float64(2*k)*math.Pi
if math.Abs(hit.Time-want) > 1e-8 {
t.Fatalf("hit %d at %.12g, want %.12g", k, hit.Time, want)
}
}
// Without the filter every half period fires: ten crossings.
hits, _, err = IntegrateODEEvents(f, 0, 5*2*math.Pi, mustFloats(t, []float64{1, 0}),
[]ODEWatch{{Function: position}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents unfiltered: %v", err)
}
if len(hits) != 10 {
t.Fatalf("unfiltered hits = %d, want 10", len(hits))
}
for i := 1; i < len(hits); i++ {
if hits[i].Time <= hits[i-1].Time {
t.Fatal("hits are not sorted by time")
}
}
}
// TestIntegrateODEStepsErrors pins the error contract of the step
// recorder.
func TestIntegrateODEStepsErrors(t *testing.T) {
y0 := mustFloats(t, []float64{1})
if _, _, err := IntegrateODESteps(decay, 0, 1, y0, ODEOptions{MaxSteps: 2}); err == nil {
t.Fatal("expected an error for an exhausted step budget")
}
if _, _, err := IntegrateODESteps(decay, 0, 1, mustFloats(t, nil), ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
}
// TestIntegrateODEStepsForward records the accepted steps of a decay:
// the trace starts at the initial state, ends exactly at t1, and every
// node sits on the analytic curve at solver tolerance.
func TestIntegrateODEStepsForward(t *testing.T) {
times, states, err := IntegrateODESteps(decay, 0, 1, mustFloats(t, []float64{1}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODESteps: %v", err)
}
if len(times) != len(states) || len(times) < 3 {
t.Fatalf("trace has %d times and %d states, want matching lengths of at least 3",
len(times), len(states))
}
if times[0] != 0 || times[len(times)-1] != 1 {
t.Fatalf("trace spans [%g, %g], want [0, 1]", times[0], times[len(times)-1])
}
for i := 1; i < len(times); i++ {
if times[i] <= times[i-1] {
t.Fatalf("times not strictly increasing at %d: %g after %g", i, times[i], times[i-1])
}
}
if states[0].FloatAt(0) != 1 {
t.Fatalf("first state = %v, want the initial state 1", states[0].FloatAt(0))
}
for i := range times {
if math.Abs(states[i].FloatAt(0)-math.Exp(-times[i])) > 1e-6 {
t.Fatalf("y(%g) = %.14g, want %.14g", times[i], states[i].FloatAt(0), math.Exp(-times[i]))
}
}
}
// TestIntegrateODEStepsDegenerate covers the zero-span trace and a
// backward pass with descending times.
func TestIntegrateODEStepsDegenerate(t *testing.T) {
times, states, err := IntegrateODESteps(decay, 1, 1, mustFloats(t, []float64{1}), ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODESteps zero span: %v", err)
}
if len(times) != 1 || times[0] != 1 || states[0].FloatAt(0) != 1 {
t.Fatalf("zero-span trace = %v, want the single initial node", times)
}
times, states, err = IntegrateODESteps(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODESteps backward: %v", err)
}
if times[0] != 1 || times[len(times)-1] != 0 {
t.Fatalf("backward trace spans [%g, %g], want [1, 0]", times[0], times[len(times)-1])
}
for i := 1; i < len(times); i++ {
if times[i] >= times[i-1] {
t.Fatalf("backward times not strictly decreasing at %d", i)
}
}
}
// TestIntegrateODEEventsErrors pins the validation and error paths.
func TestIntegrateODEEventsErrors(t *testing.T) {
f := decay
if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}),
nil, ODEOptions{}); err == nil {
t.Fatal("no watches: want an error")
}
if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}),
[]ODEWatch{{}}, ODEOptions{}); err == nil {
t.Fatal("empty watch: want an error")
}
boom := func(t float64, y *core.Array) (float64, error) {
return 0, base.Errf("watch failed")
}
if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}),
[]ODEWatch{{Function: boom}}, ODEOptions{}); err == nil {
t.Fatal("watch error: want an error")
}
// An f that fails on the refinement path must surface.
broken := func(t float64, y *core.Array) (*core.Array, error) {
if t > 0.5 {
return nil, base.Errf("integrand failed")
}
return core.MulF(y, -1), nil
}
cross := func(t float64, y *core.Array) (float64, error) {
return 1 - t, nil
}
if _, _, err := IntegrateODEEvents(broken, 0, 2, mustFloats(t, []float64{1}),
[]ODEWatch{{Function: cross}}, ODEOptions{}); err == nil {
t.Fatal("refinement error: want an error")
}
}
// TestODEArrivalGuard pins the ulp-level arrival rule: an endpoint
// missed by rounding terminates the loop, a genuine remaining distance
// does not, and a plain integration across magnitudes still lands.
func TestODEArrivalGuard(t *testing.T) {
if !odeArrived(1, 1) {
t.Fatal("equal times must count as arrived")
}
if !odeArrived(1, 1+4*2.220446049250313e-16) {
t.Fatal("a four-ulp miss must count as arrived")
}
if odeArrived(1, 1.1) {
t.Fatal("a genuine remaining distance must not count as arrived")
}
if odeArrived(0, -1e-20) {
t.Fatal("a tiny but representable distance near zero must not count as arrived")
}
// An integration whose endpoints differ well beyond Sterbenz still
// terminates and reports the endpoint state.
f := func(tt float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(0) * 0}, 1)
}
y0, _ := core.FromFloats([]float64{3}, 1)
got, err := IntegrateODE(f, 1e9, 1e9+0.5, y0, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODE across magnitudes: %v", err)
}
if got.FloatAt(0) != 3 {
t.Fatalf("constant state changed: %g", got.FloatAt(0))
}
}
+260
View File
@@ -0,0 +1,260 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"errors"
"math"
)
// The stiff workhorse beside IntegrateBackwardEuler: the variable-step
// second-order backward differentiation formula. Where the explicit
// Dormand-Prince pair must keep h·λ inside its stability region, BDF2
// is A-stable and damps the stiff mode like (2hλ)^(−1/2), so the step
// size follows accuracy alone. Each step solves the implicit relation
// α·z − h·f(t_{n+1}, z) = β by Newton over a numerical Jacobian and
// the library's LU solver, the machinery IntegrateBackwardEuler
// already carries.
//
// The step size is driven by a Milne-type estimate of the one-step
// error: the gap between the corrector and the quadratic predictor
// through the three most recent states, scaled by the constant that
// turns that gap into the BDF2 truncation error (2/11 for equal
// steps). The first step runs backward Euler, whose size a
// Hairer-Nørsett-Wanner style probe picks so the starter's own error
// already sits below the tolerance; the second BDF2 step repeats that
// size untested, safe because its truncation error is an order in h
// below the starter's; from the third step on the estimate controls
// everything.
// IntegrateBDF2 integrates y' = f(t, y) from t0 to t1 with the
// variable-step BDF2 scheme and returns y(t1). Backward integration
// works: a t1 < t0 simply integrates in the negative direction. An
// exhausted step budget, a collapsed step size, an f that returns a
// wrongly shaped state, or a Newton iteration that cannot converge
// even as the step shrinks is an error, never a silently truncated
// trajectory.
func IntegrateBDF2(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, opts ODEOptions) (*core.Array, error) {
const name = "IntegrateBDF2"
y, err := odeCheck("IntegrateBDF2", y0, &opts)
if err != nil {
return nil, err
}
n := len(y)
h, err := bdf2InitialStep(name, f, t0, t1, y, &opts)
if err != nil {
return nil, err
}
t := t0
yn := cloneDenseSlice(y)
var yNm1, yNm2 []float64
var tNm1, tNm2 float64
budget := odeBudget{max: opts.MaxSteps}
// The implicit relation's right side and the Newton seed live in
// reused buffers: both are fully rewritten at the top of every step
// and neither outlives the step's solve. The converged state lands
// in a four-buffer ring: at every acceptance the live history is
// the three most recent ring slots, so the next slot aliases
// nothing the step reads, and a rejected or stalled attempt
// reuses the slot it already holds.
beta := make([]float64, n)
seed := make([]float64, n)
var ring [4][]float64
next := 0
w := &odeWork{}
for !odeArrived(t, t1) {
if err := budget.spend(name, t, t1); err != nil {
return nil, err
}
// Never step past t1; t1−t carries the integration direction.
h = odeClampStep(h, t, t1)
tNext := t + h
hN := tNext - t
// The implicit equation and the Newton seed. The very first
// step has no history and runs backward Euler, seeded with
// the semi-implicit prediction; from the second step on the
// BDF2 weights carry the 1/h scaling themselves, so the
// derivative enters the implicit equation with weight 1.
var alpha, weight float64
estimated := yNm2 != nil
if yNm1 == nil {
alpha, weight = 1, hN
copy(beta, yn)
fy, ferr := odeEval(name, f, tNext, yn, n, &w.views)
if ferr != nil {
return nil, ferr
}
for i := range n {
seed[i] = yn[i] + hN*fy[i]
}
} else {
weight = 1
alpha = bdf2Coefficients(t, tNext, tNm1, tNm2, yn, yNm1, yNm2, beta, seed)
}
dst := ring[next]
if dst == nil {
dst = make([]float64, n)
ring[next] = dst
}
if nerr := odeNewton(name, f, w, tNext, alpha, weight, beta, seed, dst, opts.AbsTol, opts.RelTol); nerr != nil {
if errors.Is(nerr, errNewtonStalled) {
// The implicit solve struggled: halve the step and
// retry the same interval, within the step budget.
h *= 0.5
continue
}
return nil, nerr
}
factor := 1.0
if estimated {
// Milne-type local error estimate against the mixed
// absolute and relative tolerance.
c := bdf2Milne(t, tNext, tNm1, tNm2)
errNorm := 0.0
for i := range n {
scale := opts.AbsTol + opts.RelTol*math.Max(math.Abs(yn[i]), math.Abs(dst[i]))
ratio := c * (dst[i] - seed[i]) / scale
errNorm += ratio * ratio
}
errNorm = math.Sqrt(errNorm/float64(n)) + 1e-10
if errNorm <= 1 {
factor = math.Min(2, math.Max(0.2, 0.9*math.Pow(1/errNorm, 1.0/3)))
} else {
// Rejected: retry the same interval with a smaller step.
h *= math.Max(0.1, math.Min(1, 0.9*math.Pow(1/errNorm, 1.0/3)))
continue
}
}
// Accepted: shift the history one step forward. The ring slot
// becomes the working state and the buffers flow without
// copying; nothing aliases them afterwards.
yNm2, tNm2 = yNm1, tNm1
yNm1, tNm1 = yn, t
yn = dst
next = (next + 1) % len(ring)
prevT := t
t = tNext
h *= factor
// Collapse is "t did not move": a span below the absolute time
// scale is integrable, and an accepted step that arrives at the
// end exactly is not a failure either.
if t == prevT {
return nil, base.Errf("%s: the step size shrank below the resolution of t at t=%g", name, prevT)
}
}
return arrayFromVector(yn), nil
}
// bdf2InitialStep picks the first step by probing f: a trial step h0
// compares the derivative at y against the derivative one h0 further,
// and the result sizes the step so a first-order scheme's local error
// sits a factor hundred below the mixed tolerance. The span and a
// hundredfold h0 bound the answer, and the sign carries the
// integration direction.
func bdf2InitialStep(name string, f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y []float64, opts *ODEOptions) (float64, error) {
span := math.Abs(t1 - t0)
if span == 0 {
return 0, nil
}
n := len(y)
f0, err := odeEval(name, f, t0, y, n, nil)
if err != nil {
return 0, err
}
scale := make([]float64, n)
d0, d1 := 0.0, 0.0
for i := range n {
scale[i] = opts.AbsTol + opts.RelTol*math.Abs(y[i])
d0 = math.Max(d0, math.Abs(y[i])/scale[i])
d1 = math.Max(d1, math.Abs(f0[i])/scale[i])
}
h0 := 1e-6
if d0 > 1e-5 && d1 > 1e-5 {
h0 = 0.01 * d0 / d1
}
h0 = math.Min(h0, span)
// The probe steps in the integration direction: a backward span
// samples t0−h0 with the derivative subtracted, or the difference
// f1−f0 measures the wrong side of the dynamics.
dir := 1.0
if t1 < t0 {
dir = -1
}
probe := make([]float64, n)
for i := range n {
probe[i] = y[i] + dir*h0*f0[i]
}
f1, err := odeEval(name, f, t0+dir*h0, probe, n, nil)
if err != nil {
return 0, err
}
d2 := 0.0
for i := range n {
d2 = math.Max(d2, math.Abs(f1[i]-f0[i])/(scale[i]*h0))
}
h1 := span
if d := math.Max(d1, d2); d > 1e-15 {
h1 = math.Sqrt(0.01 / d)
}
h1 = math.Min(h1, math.Min(100*h0, span))
if t1 < t0 {
h1 = -h1
}
return h1, nil
}
// bdf2Coefficients assembles the variable-step BDF2 relation for a
// step from t, whose previous point sits at tNm1 (and the one before
// that at tNm2 when known), to tNext. It writes α's companions β and
// the Newton seed into the caller's buffers and returns α: the
// implicit equation is α·z − h·f(tNext, z) = β. Both buffers are fully
// overwritten. The seed is the quadratic predictor through the last
// three states when yNm2 is given, otherwise the linear ramp over the
// last two. All differences are signed, so backward integration needs
// no separate path.
func bdf2Coefficients(t, tNext, tNm1, tNm2 float64,
yn, yNm1, yNm2, beta, seed []float64) float64 {
n := len(yn)
hN := tNext - t
hP := t - tNm1
alpha := (hP + 2*hN) / ((hP + hN) * hN)
w1 := (hP + hN) / (hP * hN)
w0 := hN / (hP * (hP + hN))
for i := range n {
beta[i] = w1*yn[i] - w0*yNm1[i]
}
if yNm2 != nil {
l2 := (tNext - tNm1) * (tNext - t) / ((tNm2 - tNm1) * (tNm2 - t))
l1 := (tNext - tNm2) * (tNext - t) / ((tNm1 - tNm2) * (tNm1 - t))
l0 := (tNext - tNm2) * (tNext - tNm1) / ((t - tNm2) * (t - tNm1))
for i := range n {
seed[i] = l2*yNm2[i] + l1*yNm1[i] + l0*yn[i]
}
} else {
ramp := hN / hP
for i := range n {
seed[i] = yn[i] + ramp*(yn[i]-yNm1[i])
}
}
return alpha
}
// bdf2Milne returns the constant that turns the gap between the BDF2
// corrector and the quadratic predictor through the three previous
// states into an estimate of the corrector's one-step error: 2/11 for
// equal steps, from the leading error terms h³y”' of the predictor
// and (2/9)h³y”' of the corrector.
func bdf2Milne(t, tNext, tNm1, tNm2 float64) float64 {
hN := tNext - t
hP := t - tNm1
hPp := tNm1 - tNm2
return hN * (hN + hP) / ((hP+2*hN)*(hPp+hP+hN) + hN*(hN+hP))
}
+202
View File
@@ -0,0 +1,202 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"math"
"strings"
"testing"
)
// stiffCosine returns f for y' = −k(y − cos t), the canonical stiff
// problem: a slow forcing with a transient decaying at rate k.
func stiffCosine(k float64) func(t float64, y *core.Array) (*core.Array, error) {
return func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{-k * (y.FloatAt(0) - math.Cos(t))}, 1)
}
}
// TestIntegrateBDF2Stiff is the demonstration the stiff solver exists
// for: y' = −10^5(y − cos t) carries a transient of width 10^−5 under
// a slow forcing, and BDF2 crosses it and follows the forcing to t=1
// inside a 2000-step budget, landing on the exact solution
// y(1) = (k²·cos 1 + k·sin 1)/(k² + 1).
func TestIntegrateBDF2Stiff(t *testing.T) {
const k = 1e5
end, err := IntegrateBDF2(stiffCosine(k), 0, 1, mustFloats(t, []float64{0}),
ODEOptions{MaxSteps: 2000})
if err != nil {
t.Fatalf("IntegrateBDF2: %v", err)
}
want := (k*k*math.Cos(1) + k*math.Sin(1)) / (k*k + 1)
if math.Abs(end.FloatAt(0)-want) > 1e-6 {
t.Fatalf("y(1) = %.14g, want %.14g", end.FloatAt(0), want)
}
}
// TestIntegrateBDF2StiffBeatsExplicit shows the same problem is out
// of reach for the explicit pair: stability pins DOPRI to steps of
// order 1/k, so a 5000-step budget dies a fifth of the way in.
func TestIntegrateBDF2StiffBeatsExplicit(t *testing.T) {
_, err := IntegrateODE(stiffCosine(1e5), 0, 1, mustFloats(t, []float64{0}),
ODEOptions{MaxSteps: 5000})
if err == nil {
t.Fatal("explicit DOPRI was expected to exhaust its step budget on the stiff problem")
}
if !strings.Contains(err.Error(), "MaxSteps=5000") {
t.Fatalf("want a step-budget error, got %v", err)
}
}
// TestBDF2FixedStepOrder verifies the second order of the underlying
// formula directly: with exact history on y' = −y and uniform steps,
// halving h must quarter the global error. The Milne constant for
// equal steps is pinned to 2/11 along the way.
func TestBDF2FixedStepOrder(t *testing.T) {
if got := bdf2Milne(1, 2, 0, -1); math.Abs(got-2.0/11) > 1e-12 {
t.Fatalf("bdf2Milne for equal steps = %.14g, want 2/11", got)
}
errAt := func(steps int) float64 {
h := 1.0 / float64(steps)
alpha := 1.5 / h
now := 0.0
yn := []float64{1}
yNm1 := []float64{math.Exp(h)} // exact history at t−h
w := &odeWork{}
for range steps {
tNext := now + h
beta := []float64{2*yn[0]/h - yNm1[0]/(2*h)}
z := make([]float64, 1)
err := odeNewton("TestBDF2FixedStepOrder", decay, w, tNext, alpha, 1,
beta, []float64{math.Exp(-tNext)}, z, 1e-13, 1e-13)
if err != nil {
t.Fatalf("odeNewton: %v", err)
}
yNm1 = yn
yn = z
now = tNext
}
return math.Abs(yn[0] - math.Exp(-1))
}
e20, e40 := errAt(20), errAt(40)
if e20 < 1e-12 {
t.Skipf("error already at round-off (%v)", e20)
}
ratio := e20 / e40
if ratio < 3 || ratio > 5.2 {
t.Fatalf("error ratio over a halved step = %.2g, want ≈ 4 for a second-order scheme", ratio)
}
}
// TestIntegrateBDF2Accuracy checks the adaptive driver on a smooth
// problem against the analytic decay at a tolerance far below the
// default, and over a full oscillator period with a two-dimensional
// state, exercising the vector Newton path.
func TestIntegrateBDF2Accuracy(t *testing.T) {
end, err := IntegrateBDF2(decay, 0, 1, mustFloats(t, []float64{1}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateBDF2: %v", err)
}
if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-5 {
t.Fatalf("y(1) = %.14g, want %.14g ± 1e-5", end.FloatAt(0), math.Exp(-1))
}
oscillator := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
full, err := IntegrateBDF2(oscillator, 0, 2*math.Pi, mustFloats(t, []float64{1, 0}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateBDF2 oscillator: %v", err)
}
if math.Abs(full.FloatAt(0)-1) > 1e-4 || math.Abs(full.FloatAt(1)) > 1e-4 {
t.Fatalf("full period = (%.10g, %.10g), want (1, 0)",
full.FloatAt(0), full.FloatAt(1))
}
}
// TestIntegrateBDF2Backward integrates the decay backwards from t=1
// to t=0; the signed-step formulation must return the start value.
func TestIntegrateBDF2Backward(t *testing.T) {
end, err := IntegrateBDF2(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateBDF2 backward: %v", err)
}
if math.Abs(end.FloatAt(0)-1) > 1e-5 {
t.Fatalf("backward y(0) = %.14g, want 1 ± 1e-5", end.FloatAt(0))
}
}
// TestIntegrateBDF2Errors pins the error contract: a degenerate span
// returns the initial state unchanged, a wrong-shaped f, a rank-2
// state, an empty state and an exhausted step budget are errors.
func TestIntegrateBDF2Errors(t *testing.T) {
y0 := mustFloats(t, []float64{1})
same, err := IntegrateBDF2(decay, 1, 1, y0, ODEOptions{})
if err != nil {
t.Fatalf("zero span: %v", err)
}
if math.Abs(same.FloatAt(0)-1) > 0 {
t.Fatalf("zero span moved the state to %v", same.FloatAt(0))
}
wrongShape := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{1, 1}, 2)
}
if _, err := IntegrateBDF2(wrongShape, 0, 1, y0, ODEOptions{}); err == nil {
t.Fatal("expected an error when f returns the wrong shape")
}
matrixState := mustFloats(t, []float64{1, 1}, 1, 2)
if _, err := IntegrateBDF2(decay, 0, 1, matrixState, ODEOptions{}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
if _, err := IntegrateBDF2(decay, 0, 1, mustFloats(t, nil), ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
if _, err := IntegrateBDF2(decay, 0, 1, y0, ODEOptions{MaxSteps: 2}); err == nil {
t.Fatal("expected an error for an exhausted step budget")
}
boom := func(t float64, y *core.Array) (*core.Array, error) {
if t > 0.5 {
return nil, base.Errf("detector tripped")
}
return core.MulF(y, -1), nil
}
if _, err := IntegrateBDF2(boom, 0, 1, y0, ODEOptions{}); err == nil {
t.Fatal("expected the operator error to propagate")
}
}
// TestBDF2InitialStepBackwardProbe pins the probe direction: on a
// backward span the initial-step probe must sample the dynamics at
// t0 − h0, not extrapolate forward, so every evaluation time either
// sits before t0 or the probe is wrong.
func TestBDF2InitialStepBackwardProbe(t *testing.T) {
var calls []float64
f := func(tt float64, y *core.Array) (*core.Array, error) {
calls = append(calls, tt)
return mustFloats(t, []float64{0}), nil
}
h, err := bdf2InitialStep("TestBDF2InitialStep", f, 5, 0, []float64{1}, &ODEOptions{})
if err != nil {
t.Fatalf("bdf2InitialStep: %v", err)
}
if h >= 0 {
t.Fatalf("backward span must yield a negative first step, got %g", h)
}
backward := false
for _, c := range calls {
if c < 5 {
backward = true
}
}
if !backward {
t.Fatalf("the probe never stepped backward from t0 = 5, evaluated at %v", calls)
}
}
+414
View File
@@ -0,0 +1,414 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"errors"
"math"
)
// The variable-order stiff workhorse above IntegrateBDF2: the backward
// differentiation formula of order one through five, with the step size
// and the order both adapted every step in the VODE manner. Each step
// interpolates a polynomial of degree k through the k most recent
// states and the unknown end value and requires its derivative at the
// new time to equal f, the same implicit relation BDF2 solves; the
// Newton iteration, the LU machinery, the Hairer-Nørsett-Wanner initial
// step probe and the step controller are the ones IntegrateBDF2
// already carries.
//
// The coefficients are the variable-step, divided-difference form: the
// Newton form of the interpolating polynomial through (tNext, z) and
// the stored back values, written per component from a small divided-
// difference table over the stored times. The form was chosen over the
// fixed-coefficient one because the package keeps a solution history
// rather than a Nordsieck array, because the divided differences feed
// the order selection (the a-priori error estimate per candidate order
// falls out of the same table) and because the relation leaves the
// Newton contract α·z − h·f(tNext, z) = β of odeNewton untouched. At
// order two with equal steps the assembled α and β agree with
// bdf2Coefficients to rounding, so the shipped BDF2 behaviour is the
// special case the driver degrades to.
//
// The local error estimate is the Milne-type one: the gap between the
// corrector and the degree-k predictor extrapolated from the k+1
// newest states, scaled by the constant that turns the gap into the
// corrector's own error. The variable-step constant generalises the
// 2/11 of bdf2Milne: with α the derivative weight of the new point and
// S the span from tNext to the oldest predictor node, the estimate is
// (z − seed)/(1 + α·S), which for equal steps of order two reproduces
// 2/11 exactly. The order itself is chosen before the solve, from the
// divided differences of the stored states: the (k+1)-th divided
// difference approximates y^(k+1)/(k+1)!, and the candidate whose
// implied optimal step is largest wins, with a margin so the order
// does not flicker between neighbours.
//
// The first step is backward Euler, sized by the shared probe; the
// order ramps up as the history accumulates, one level per step.
// BDFVarStats reports what a variable-order run did: the accepted and
// rejected steps and the highest order the driver reached.
type BDFVarStats struct {
Steps int
Rejected int
MaxOrder int
}
// BDFVarOptions tunes IntegrateBDFVar. RelTol ≤ 0 means 1e-6, AbsTol ≤ 0
// means 1e-9, MaxSteps ≤ 0 means 100000, the ODEOptions defaults. Stats,
// when not nil, receives the run's counters.
type BDFVarOptions struct {
RelTol float64
AbsTol float64
MaxSteps int
Stats *BDFVarStats
}
// bdfVarOrderMax is the highest order the driver raises to. bdfVarKeep
// is the number of states held back: order k needs k back values for
// its corrector, k+1 for its predictor and k+2 for the a-priori order
// comparison, so seven states serve order five in every role.
const (
bdfVarOrderMax = 5
bdfVarKeep = bdfVarOrderMax + 2
)
// IntegrateBDFVar integrates y' = f(t, y) from t0 to t1 with the
// variable-step, variable-order BDF scheme of orders one through five
// and returns y(t1). Backward integration works: a t1 < t0 simply
// integrates in the negative direction. An exhausted step budget, a
// collapsed step size, an f that returns a wrongly shaped state, or a
// Newton iteration that cannot converge even as the step shrinks is an
// error, never a silently truncated trajectory.
func IntegrateBDFVar(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, opts BDFVarOptions) (*core.Array, error) {
end, err := integrateBDFVar("IntegrateBDFVar", f, t0, t1, y0, opts, bdfVarOrderMax, false)
if err != nil {
return nil, err
}
return arrayFromVector(end), nil
}
// integrateBDFVar drives the variable-order loop. maxOrder caps the
// order adaptation and lockOrder pins the order at maxOrder once the
// history ramp reaches it, which is the fixed-order hook the tests
// drive; the public entry always asks for adaptive order five.
func integrateBDFVar(name string, f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, opts BDFVarOptions, maxOrder int, lockOrder bool) ([]float64, error) {
if maxOrder < 1 || maxOrder > bdfVarOrderMax {
return nil, base.Errf("%s: maxOrder must be between 1 and %d, got %d", name, bdfVarOrderMax, maxOrder)
}
y, err := odeCheck(name, y0, nil)
if err != nil {
return nil, err
}
relTol, absTol, maxSteps := opts.RelTol, opts.AbsTol, opts.MaxSteps
if relTol <= 0 {
relTol = 1e-6
}
if absTol <= 0 {
absTol = 1e-9
}
if maxSteps <= 0 {
maxSteps = 100000
}
n := len(y)
h, err := bdf2InitialStep(name, f, t0, t1, y, &ODEOptions{RelTol: relTol, AbsTol: absTol})
if err != nil {
return nil, err
}
var stats BDFVarStats
hist := &bdfVarHistory{}
hist.push(t0, y)
// The implicit relation's right side, the predictor seed and the
// divided-difference workspace live in reused buffers: all are fully
// rewritten at the top of every step. The Newton result lands in a
// per-solve scratch buffer that never touches the history window:
// the order selection and the coefficient assembly reread the whole
// window on a retry, so a rejected attempt must leave every stored
// state intact. Only an accepted step copies the state into the
// ring slot its push then occupies.
beta := make([]float64, n)
seed := make([]float64, n)
zbuf := make([]float64, n)
dd := make([]float64, bdfVarKeep)
nodes := make([]float64, bdfVarKeep)
spans := make([]float64, bdfVarKeep)
ddTab := make([][]float64, bdfVarKeep)
for level := range ddTab {
ddTab[level] = make([]float64, n)
}
budget := odeBudget{max: maxSteps}
w := &odeWork{}
t := t0
carried := 1
for !odeArrived(t, t1) {
if err := budget.spend(name, t, t1); err != nil {
return nil, err
}
// Never step past t1; t1−t carries the integration direction.
h = odeClampStep(h, t, t1)
tNext := t + h
hN := tNext - t
var alpha, weight float64
var order int
estimated := false
if hist.n == 1 {
// The very first step has no history and runs backward
// Euler, seeded with the semi-implicit prediction: the
// house starter IntegrateBDF2 begins with.
alpha, weight, order = 1, hN, 1
copy(beta, y)
fy, ferr := odeEval(name, f, tNext, y, n, &w.views)
if ferr != nil {
return nil, ferr
}
for i := range n {
seed[i] = y[i] + hN*fy[i]
}
} else {
order = min(carried, maxOrder, hist.n-1)
switch {
case lockOrder && hist.n > maxOrder:
// The fixed-order contract: once the history ramp can
// feed the requested order, every step runs at it.
order = maxOrder
case !lockOrder && hist.n >= 3:
bdfVarDividedDifferences(hist, n, ddTab, dd, nodes)
order = bdfVarPickOrder(order, maxOrder, hist.n, tNext, h, hist, y, absTol, relTol, ddTab, spans)
}
weight = 1
alpha = bdfVarCoefficients(order, tNext, hist, beta, seed, dd, nodes, spans)
estimated = true
}
if nerr := odeNewton(name, f, w, tNext, alpha, weight, beta, seed, zbuf, absTol, relTol); nerr != nil {
if errors.Is(nerr, errNewtonStalled) {
// The implicit solve struggled: halve the step and
// retry the same interval, within the step budget.
h *= 0.5
continue
}
return nil, nerr
}
factor := 1.0
if estimated {
// Milne-type local error estimate against the mixed
// absolute and relative tolerance.
_, tOldest := hist.back(order)
c := 1 / (1 + alpha*(tNext-tOldest))
errNorm := 0.0
for i := range n {
scale := absTol + relTol*math.Max(math.Abs(y[i]), math.Abs(zbuf[i]))
ratio := c * (zbuf[i] - seed[i]) / scale
errNorm += ratio * ratio
}
errNorm = math.Sqrt(errNorm/float64(n)) + 1e-10
if errNorm <= 1 {
factor = min(2, max(0.2, 0.9*math.Pow(1/errNorm, 1/float64(order+1))))
} else {
// Rejected: retry the same interval with a smaller step.
stats.Rejected++
h *= max(0.1, min(1, 0.9*math.Pow(1/errNorm, 1/float64(order+1))))
continue
}
}
// Accepted: the corrector is copied into the ring slot the push
// fills and becomes the working state, so back(0) is always
// (t, y) and the buffers flow without copying.
slot := hist.y[hist.next]
if slot == nil {
slot = make([]float64, n)
}
copy(slot, zbuf)
hist.push(tNext, slot)
y = slot
carried = order
if order > stats.MaxOrder {
stats.MaxOrder = order
}
stats.Steps++
prevT := t
t = tNext
h *= factor
// Collapse is "t did not move": a span below the absolute time
// scale is integrable, and an accepted step that arrives at the
// end exactly is not a failure either.
if t == prevT {
return nil, base.Errf("%s: the step size shrank below the resolution of t at t=%g", name, prevT)
}
}
if opts.Stats != nil {
*opts.Stats = stats
}
return y, nil
}
// bdfVarHistory holds the last bdfVarKeep accepted states with their
// times in a fixed ring. back(0) is the newest state, back(1) the one
// before it, and so on; slots are recycled only once they are too old
// to serve any order, so the buffers flow without copying.
type bdfVarHistory struct {
y [bdfVarKeep][]float64
t [bdfVarKeep]float64
next int
n int
}
// push records an accepted state and its time as the new newest entry.
func (h *bdfVarHistory) push(t float64, y []float64) {
h.y[h.next], h.t[h.next] = y, t
h.next = (h.next + 1) % bdfVarKeep
if h.n < bdfVarKeep {
h.n++
}
}
// back returns the state i steps behind the newest one.
func (h *bdfVarHistory) back(i int) ([]float64, float64) {
j := (h.next - 1 - i + bdfVarKeep) % bdfVarKeep
return h.y[j], h.t[j]
}
// bdfVarCoefficients assembles the variable-step BDF relation of the
// given order for a step to tNext from the newest history state. It
// writes α's companions β and the Newton seed into the caller's
// buffers and returns α: the implicit equation is α·z − f(tNext, z) =
// β, the weight already scaled out. Both buffers are fully overwritten.
// The seed is the degree-order polynomial through the order+1 newest
// states evaluated at tNext, the predictor the error estimate reads.
// All differences are signed, so backward integration needs no
// separate path.
//
// The construction is the divided-difference (Newton) form: with nodes
// x_0 = tNext and x_q = the q-th back time, the interpolating
// polynomial's derivative at tNext is Σ_j c_j·Π_j where c_j are the
// divided differences of the data (z at x_0, the back values after)
// and Π_j the Newton basis products. Splitting c_j into its z part,
// 1/Π_j, and its history part gives α = Σ 1/(tNext − x_m), the Lagrange
// derivative weight of the new point, and β from the history-only
// table, all from one per-component recursion.
func bdfVarCoefficients(order int, tNext float64, hist *bdfVarHistory,
beta, seed, dd, nodes, spans []float64) float64 {
nodes[0] = tNext
// The ring's nodes and value slices are the same for every
// component: gather both once, on the stack, instead of walking the
// ring inside the per-element loop.
var backVals [bdfVarKeep][]float64
for q := range order + 1 {
backVals[q], nodes[q+1] = hist.back(q)
}
// spans[m] is Π_m, the product of tNext − x_q over q < m: the
// Newton basis value the level-m coefficients multiply.
spans[0] = 1
for m := 1; m <= order; m++ {
spans[m] = spans[m-1] * (tNext - nodes[m])
}
alpha := 0.0
for m := 1; m <= order; m++ {
alpha += 1 / (tNext - nodes[m])
}
for i := range beta {
// dd[q] starts as the value at node q: zero at tNext, the back
// values after. One level of the recursion per Newton term;
// level order leaves dd[0] holding the order-th divided
// difference over the new point and dd[1] the one over the
// stored values, which is the predictor's top coefficient.
dd[0] = 0
for q := range order + 1 {
dd[q+1] = backVals[q][i]
}
seed[i] = dd[1]
betaSum := 0.0
for level := 1; level <= order; level++ {
for q := range order + 2 - level {
dd[q] = (dd[q+1] - dd[q]) / (nodes[q+level] - nodes[q])
}
betaSum += dd[0] * spans[level-1]
seed[i] += dd[1] * spans[level]
}
beta[i] = -betaSum
}
return alpha
}
// bdfVarDividedDifferences fills tab with the divided differences of
// the stored back values alone: tab[level][i] is the level-th divided
// difference of (y_n, y_{n-1}, …) over their times for component i.
// The (order+1)-th entry approximates y^(order+1)/(order+1)! and is
// what the a-priori order comparison reads.
func bdfVarDividedDifferences(hist *bdfVarHistory, n int, tab [][]float64, dd, times []float64) {
// The ring's times and value slices do not depend on the component:
// gather both once, on the stack, instead of walking the ring
// inside the per-element loops.
var backVals [bdfVarKeep][]float64
for q := range hist.n {
backVals[q], times[q] = hist.back(q)
}
for i := range n {
for q := range hist.n {
dd[q] = backVals[q][i]
}
for level := 1; level < hist.n; level++ {
for q := range hist.n - level {
dd[q] = (dd[q+1] - dd[q]) / (times[q+level] - times[q])
}
tab[level][i] = dd[0]
}
}
}
// bdfVarPickOrder returns the order for the coming step. Every order
// the history supports gets an a-priori optimal step: the local error
// the divided differences predict, raised to the power that would
// bring it to the tolerance. The scan runs from order 1 upward and a
// candidate must beat the running best by a clear margin, so the
// effective pick is the lowest order within 15 percent of the largest
// predicted step: short histories and cheap coefficients win near
// ties, and the order settles instead of flickering between equals.
// The carried order survives the scan only before a second state is
// held; after that some candidate always displaces it.
func bdfVarPickOrder(carried, maxOrder, held int, tNext, h float64, hist *bdfVarHistory, y []float64,
absTol, relTol float64, tab [][]float64, spans []float64) int {
best, bestH := carried, 0.0
for j := 1; j <= min(maxOrder, held-1); j++ {
hj := math.Abs(h)
if j <= held-2 {
e := bdfVarPriorNorm(j, tNext, hist, y, absTol, relTol, tab, spans)
hj = math.Abs(h) * math.Pow(1/e, 1/float64(j+1))
}
if hj > bestH*1.15 {
best, bestH = j, hj
}
}
return best
}
// bdfVarPriorNorm estimates the RMS error norm a step of size h at the
// given order would produce: the (order+1)-th divided difference of
// the stored states approximates y^(order+1)/(order+1)!, and the
// order's local error scales that by the Newton basis product over α,
// the same estimate the Milne constant formalises a posteriori.
func bdfVarPriorNorm(order int, tNext float64, hist *bdfVarHistory, y []float64,
absTol, relTol float64, tab [][]float64, spans []float64) float64 {
spans[0] = 1
alpha := 0.0
for q := range order {
_, tq := hist.back(q)
spans[q+1] = spans[q] * math.Abs(tNext-tq)
alpha += 1 / math.Abs(tNext-tq)
}
w := spans[order] / alpha
norm := 0.0
for i := range y {
scale := absTol + relTol*math.Abs(y[i])
ratio := math.Abs(tab[order+1][i]) * w / scale
norm += ratio * ratio
}
return math.Sqrt(norm/float64(len(y))) + 1e-10
}
+298
View File
@@ -0,0 +1,298 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"math"
"testing"
)
// TestIntegrateBDFVarStiff is the demonstration pin: on y' =
// −10^5(y − cos t) the variable-order driver lands on the exact y(1) =
// (k²·cos 1 + k·sin 1)/(k² + 1) inside half the step budget BDF2
// needed, having raised to order five on the smooth tail.
func TestIntegrateBDFVarStiff(t *testing.T) {
const k = 1e5
var stats BDFVarStats
end, err := IntegrateBDFVar(stiffCosine(k), 0, 1, mustFloats(t, []float64{0}),
BDFVarOptions{MaxSteps: 2000, Stats: &stats})
if err != nil {
t.Fatalf("IntegrateBDFVar: %v", err)
}
want := (k*k*math.Cos(1) + k*math.Sin(1)) / (k*k + 1)
if math.Abs(end.FloatAt(0)-want) > 1e-6 {
t.Fatalf("y(1) = %.14g, want %.14g", end.FloatAt(0), want)
}
t.Logf("stiff run: %d steps, %d rejected, max order %d", stats.Steps, stats.Rejected, stats.MaxOrder)
if stats.MaxOrder != 5 {
t.Fatalf("max order reached = %d, want 5 on the smooth tail", stats.MaxOrder)
}
if stats.Steps > 1000 {
t.Fatalf("the run took %d steps, want well inside the 2000-step budget BDF2 needed", stats.Steps)
}
}
// TestIntegrateBDFVarOrderAdapts instruments the order counter: the
// first accepted steps run at order one (nothing else has history), so
// a run that ends with order five must have climbed the ladder, and on
// the same stiff problem it must spend far fewer steps than an
// order-one-locked run, which is what step and order adaptation buy.
func TestIntegrateBDFVarOrderAdapts(t *testing.T) {
const k = 1e5
var adaptive, locked BDFVarStats
if _, err := IntegrateBDFVar(stiffCosine(k), 0, 1, mustFloats(t, []float64{0}),
BDFVarOptions{MaxSteps: 50000, Stats: &adaptive}); err != nil {
t.Fatalf("IntegrateBDFVar adaptive: %v", err)
}
if _, err := integrateBDFVar("TestIntegrateBDFVarOrderAdapts", stiffCosine(k), 0, 1,
mustFloats(t, []float64{0}), BDFVarOptions{MaxSteps: 50000, Stats: &locked}, 1, true); err != nil {
t.Fatalf("IntegrateBDFVar order-one locked: %v", err)
}
t.Logf("adaptive run: %d steps, locked run: %d steps", adaptive.Steps, locked.Steps)
if adaptive.Steps < 8 || locked.Steps < 8 {
t.Fatalf("implausible step counts: adaptive %d, locked %d", adaptive.Steps, locked.Steps)
}
if adaptive.MaxOrder != 5 {
t.Fatalf("adaptive run reached order %d, want 5", adaptive.MaxOrder)
}
if locked.MaxOrder != 1 {
t.Fatalf("locked run reached order %d, want 1 throughout", locked.MaxOrder)
}
if adaptive.Steps*3 > locked.Steps {
t.Fatalf("the adaptive run took %d steps against the locked run's %d: order adaptation did not engage",
adaptive.Steps, locked.Steps)
}
}
// TestIntegrateBDFVarAccuracy checks the adaptive driver on a smooth
// problem against the analytic decay, over a full oscillator period
// with a two-dimensional state, and backwards in time.
func TestIntegrateBDFVarAccuracy(t *testing.T) {
end, err := IntegrateBDFVar(decay, 0, 1, mustFloats(t, []float64{1}),
BDFVarOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateBDFVar: %v", err)
}
if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-5 {
t.Fatalf("y(1) = %.14g, want %.14g ± 1e-5", end.FloatAt(0), math.Exp(-1))
}
oscillator := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
full, err := IntegrateBDFVar(oscillator, 0, 2*math.Pi, mustFloats(t, []float64{1, 0}),
BDFVarOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateBDFVar oscillator: %v", err)
}
if math.Abs(full.FloatAt(0)-1) > 1e-4 || math.Abs(full.FloatAt(1)) > 1e-4 {
t.Fatalf("full period = (%.10g, %.10g), want (1, 0)",
full.FloatAt(0), full.FloatAt(1))
}
back, err := IntegrateBDFVar(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}),
BDFVarOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateBDFVar backward: %v", err)
}
if math.Abs(back.FloatAt(0)-1) > 1e-5 {
t.Fatalf("backward y(0) = %.14g, want 1 ± 1e-5", back.FloatAt(0))
}
}
// TestBDFVarCoefficientsMatchBDF2 pins the coefficient recurrence: at
// order two the divided-difference form must reproduce the shipped
// bdf2Coefficients, on equal steps and on skewed ones, in α, β and the
// predictor seed alike.
func TestBDFVarCoefficientsMatchBDF2(t *testing.T) {
patterns := []struct{ tNext, t, tNm1, tNm2 float64 }{
{3, 2, 1, 0},
{1.3, 0.75, 0.4, -0.1},
{5, 1, 0.5, -2},
}
vals := []float64{2.5, -3, 7} // y at tNm2, tNm1, t
for _, p := range patterns {
hist := &bdfVarHistory{}
for i, tt := range []float64{p.tNm2, p.tNm1, p.t} {
hist.push(tt, []float64{vals[i]})
}
beta := make([]float64, 1)
seed := make([]float64, 1)
alpha := bdfVarCoefficients(2, p.tNext, hist, beta, seed,
make([]float64, bdfVarKeep), make([]float64, bdfVarKeep), make([]float64, bdfVarKeep))
beta2 := make([]float64, 1)
seed2 := make([]float64, 1)
alpha2 := bdf2Coefficients(p.t, p.tNext, p.tNm1, p.tNm2, []float64{vals[2]},
[]float64{vals[1]}, []float64{vals[0]}, beta2, seed2)
tol := func(v float64) float64 { return 1e-12 * math.Max(1, math.Abs(v)) }
if math.Abs(alpha-alpha2) > tol(alpha2) {
t.Fatalf("pattern %v: alpha = %.16g, bdf2 gives %.16g", p, alpha, alpha2)
}
if math.Abs(beta[0]-beta2[0]) > tol(beta2[0]) {
t.Fatalf("pattern %v: beta = %.16g, bdf2 gives %.16g", p, beta[0], beta2[0])
}
if math.Abs(seed[0]-seed2[0]) > tol(seed2[0]) {
t.Fatalf("pattern %v: seed = %.16g, bdf2 gives %.16g", p, seed[0], seed2[0])
}
}
}
// TestBDFVarMilneConstantMatchesBDF2 pins the variable-step Milne
// constant against the shipped bdf2Milne at order two.
func TestBDFVarMilneConstantMatchesBDF2(t *testing.T) {
patterns := []struct{ tNext, t, tNm1, tNm2 float64 }{
{3, 2, 1, 0},
{1.3, 0.75, 0.4, -0.1},
{5, 1, 0.5, -2},
}
for _, p := range patterns {
hist := &bdfVarHistory{}
for _, tt := range []float64{p.tNm2, p.tNm1, p.t} {
hist.push(tt, []float64{0})
}
alpha := bdfVarCoefficients(2, p.tNext, hist, make([]float64, 1), make([]float64, 1),
make([]float64, bdfVarKeep), make([]float64, bdfVarKeep), make([]float64, bdfVarKeep))
_, tOldest := hist.back(2)
got := 1 / (1 + alpha*(p.tNext-tOldest))
want := bdf2Milne(p.t, p.tNext, p.tNm1, p.tNm2)
if math.Abs(got-want) > 1e-14*math.Max(1, math.Abs(want)) {
t.Fatalf("pattern %v: milne constant %.16g, bdf2Milne gives %.16g", p, got, want)
}
}
}
// TestIntegrateBDFVarFixedOrderLinear pins the exactness of the fixed
// orders on y' = k·t^(k−1), whose solution y = t^k only order k
// reproduces exactly: the k-step formula carries the k-th derivative
// the problem is built from, and any lower order drops it, so each
// locked run must land on 1 at the end AND report that it ran at the
// locked order, which together rule out a hook that silently
// integrates at order 1.
func TestIntegrateBDFVarFixedOrderLinear(t *testing.T) {
for order := 1; order <= 5; order++ {
k := float64(order)
f := func(t float64, y *core.Array) (*core.Array, error) {
out, err := core.Zeros(core.Float, 1)
if err != nil {
return nil, err
}
out.SetFloatAt(0, k*math.Pow(t, k-1))
return out, nil
}
stats := &BDFVarStats{}
end, err := integrateBDFVar("TestIntegrateBDFVarFixedOrderLinear", f, 0, 1,
mustFloats(t, []float64{0}), BDFVarOptions{MaxSteps: 10000, Stats: stats}, order, true)
if err != nil {
t.Fatalf("locked order %d: %v", order, err)
}
if math.Abs(end[0]-1) > 5e-5 {
t.Fatalf("locked order %d: y(1) = %.16g, want 1", order, end[0])
}
if stats.MaxOrder != order {
t.Fatalf("locked order %d ran at max order %d", order, stats.MaxOrder)
}
}
}
// TestBDFVarExactPolynomialPerOrder drives the coefficient recurrence
// directly: a single order-k step from exact history on the degree-k
// polynomial p(t) = t^k must return p at the new time to rounding, on
// skewed steps, because the variable-step formula is exact for degree
// k when the past is exact.
func TestBDFVarExactPolynomialPerOrder(t *testing.T) {
const tNext = 1.3
// Back-value times on skewed step gaps, newest first.
patterns := [][6]float64{
{1, 0.7, 0.35, 0.1, -0.2, -1},
{1, 0.9, 0.75, 0.5, 0.2, -0.1},
}
for order := 1; order <= 5; order++ {
for _, g := range patterns {
times := g[:order+1]
hist := &bdfVarHistory{}
for _, tt := range times {
hist.push(tt, []float64{math.Pow(tt, float64(order))})
}
beta := make([]float64, 1)
seed := make([]float64, 1)
w := &odeWork{}
alpha := bdfVarCoefficients(order, tNext, hist, beta, seed,
make([]float64, bdfVarKeep), make([]float64, bdfVarKeep), make([]float64, bdfVarKeep))
z := make([]float64, 1)
err := odeNewton("TestBDFVarExactPolynomialPerOrder",
func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{float64(order) * math.Pow(t, float64(order-1))}, 1)
}, w, tNext, alpha, 1, beta, seed, z, 1e-13, 1e-13)
if err != nil {
t.Fatalf("order %d gaps %v: odeNewton: %v", order, times, err)
}
want := math.Pow(tNext, float64(order))
if math.Abs(z[0]-want) > 1e-11*math.Max(1, math.Abs(want)) {
t.Fatalf("order %d gaps %v: z = %.16g, want %.16g to rounding", order, times, z[0], want)
}
}
}
}
// TestIntegrateBDFVarErrors pins the error contract: a degenerate span
// returns the initial state unchanged, a wrong-shaped f, a rank-2
// state, an empty state and an exhausted step budget are errors, and a
// nonsensical order cap is refused.
func TestIntegrateBDFVarErrors(t *testing.T) {
y0 := mustFloats(t, []float64{1})
same, err := integrateBDFVar("TestIntegrateBDFVarErrors", decay, 1, 1, y0, BDFVarOptions{}, 5, false)
if err != nil {
t.Fatalf("zero span: %v", err)
}
if math.Abs(same[0]-1) > 0 {
t.Fatalf("zero span moved the state to %v", same[0])
}
wrongShape := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{1, 1}, 2)
}
if _, err := IntegrateBDFVar(wrongShape, 0, 1, y0, BDFVarOptions{}); err == nil {
t.Fatal("expected an error when f returns the wrong shape")
}
matrixState := mustFloats(t, []float64{1, 1}, 1, 2)
if _, err := IntegrateBDFVar(decay, 0, 1, matrixState, BDFVarOptions{}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
if _, err := IntegrateBDFVar(decay, 0, 1, mustFloats(t, nil), BDFVarOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
if _, err := IntegrateBDFVar(decay, 0, 1, y0, BDFVarOptions{MaxSteps: 2}); err == nil {
t.Fatal("expected an error for an exhausted step budget")
}
if _, err := integrateBDFVar("TestIntegrateBDFVarErrors", decay, 0, 1, y0, BDFVarOptions{}, 6, false); err == nil {
t.Fatal("expected an error for an order cap above five")
}
if _, err := integrateBDFVar("TestIntegrateBDFVarErrors", decay, 0, 1, y0, BDFVarOptions{}, 0, false); err == nil {
t.Fatal("expected an error for an order cap below one")
}
boom := func(t float64, y *core.Array) (*core.Array, error) {
if t > 0.5 {
return nil, base.Errf("detector tripped")
}
return core.MulF(y, -1), nil
}
if _, err := IntegrateBDFVar(boom, 0, 1, y0, BDFVarOptions{}); err == nil {
t.Fatal("expected the operator error to propagate")
}
// An f that survives the two probe evaluations and fails on the
// starter's own evaluation is refused at once.
calls := 0
counted := func(t float64, y *core.Array) (*core.Array, error) {
calls++
if calls > 2 {
return nil, base.Errf("detector tripped")
}
return core.FromFloats([]float64{0}, 1)
}
if _, err := IntegrateBDFVar(counted, 0, 1, mustFloats(t, []float64{1}), BDFVarOptions{}); err == nil {
t.Fatal("expected the starter's f failure to surface")
}
}

Some files were not shown because too many files have changed in this diff Show More