feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -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/...
|
||||
@@ -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};
|
||||
'
|
||||
@@ -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 .
|
||||
@@ -0,0 +1,5 @@
|
||||
.idea/
|
||||
.zcode/
|
||||
bin/
|
||||
coverage.out
|
||||
*.test
|
||||
+348
@@ -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
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
@@ -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.
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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] }
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+313
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user